File size: 9,351 Bytes
9d45136 3db44c2 9d45136 3db44c2 9d45136 3db44c2 9d45136 3db44c2 9d45136 3db44c2 9d45136 3db44c2 9d45136 3db44c2 9d45136 d9c9969 9d45136 d9c9969 9d45136 3db44c2 9d45136 3db44c2 9d45136 3db44c2 9d45136 3db44c2 9d45136 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 | #!/usr/bin/env python
"""
Inference for the minimal SpeechLLM -- the tutorial's demo entry point.
# one file
python inference.py --ckpt exp/run/checkpoints/last.ckpt --audio sample.flac
# score a manifest: prints reference vs hypothesis side by side
python inference.py --ckpt exp/run/checkpoints/last.ckpt \
--manifest data/Librispeech-dev-test/test-clean_asr.jsonl \
--audio-root /path/to/data --limit 20
"""
import argparse
import json
import logging
import os
import torch
from transformers import AutoFeatureExtractor, AutoTokenizer
from data import Collator
from modeling import SpeechLLM
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
logger = logging.getLogger(__name__)
DEFAULT_PROMPT = "<audio><|AUDIO|></audio>\n\nTranscribe the speech into text"
class SpeechLLMForInference:
"""Checkpoint in, text out.
`generate()` takes DeSTA3-style messages so the demo reads like a chat call:
[{"role": "user",
"content": "<audio><|AUDIO|></audio>\\n\\nTranscribe the speech into text",
"audios": [{"audio": "sample.flac"}]}]
Pass a list of those to batch several utterances in one forward pass.
"""
def __init__(self, model, tokenizer, collator, device, dtype):
self.model, self.tokenizer = model, tokenizer
self.collator, self.device, self.dtype = collator, device, dtype
# -- loading --
@classmethod
def from_checkpoint(cls, ckpt_path, device=None, dtype=None):
device = device or ("cuda" if torch.cuda.is_available() else "cpu")
ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
hp = ckpt["hyper_parameters"]
dtype = dtype or getattr(torch, hp["dtype"])
logger.info("checkpoint %s (epoch %s, step %s)",
os.path.basename(ckpt_path), ckpt.get("epoch"), ckpt.get("global_step"))
# n_downsample used to be n_downsample_layers, an exponent: 1 meant 2 frames
# per token. Map the old key so earlier checkpoints still load.
n_downsample = hp.get("n_downsample") or 2 ** hp["n_downsample_layers"]
logger.info(" llm=%s encoder=%s n_downsample=%d frozen_encoder=%s",
hp["llm_id"], hp["encoder_id"],
n_downsample, hp.get("freeze_encoder", False))
# rebuilt exactly as SpeechLLMModule.__init__ does, or the ids shift
tokenizer = AutoTokenizer.from_pretrained(hp["llm_id"])
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "left"
tokenizer.add_tokens([hp["audio_locator"]])
feature_extractor = AutoFeatureExtractor.from_pretrained(hp["encoder_id"])
model = SpeechLLM(
llm_id=hp["llm_id"], encoder_id=hp["encoder_id"],
n_downsample=n_downsample, dtype=dtype,
# .get(): older checkpoints predate these options
adapter_hidden_dim=hp.get("adapter_hidden_dim", 0),
freeze_encoder=hp.get("freeze_encoder", False), use_lora=True,
lora_rank=hp["lora_rank"], lora_alpha=hp["lora_alpha"],
lora_dropout=hp["lora_dropout"],
lora_target_modules=hp["lora_target_modules"],
)
cls._load_trainable(model, ckpt["state_dict"])
model.to(device).eval()
collator = Collator(
tokenizer=tokenizer, feature_extractor=feature_extractor,
audio_locator=hp["audio_locator"], placeholder_token=hp["placeholder_token"],
max_seq_length=hp["max_seq_length"], max_audio_seconds=hp["max_audio_seconds"],
n_downsample=n_downsample,
for_generation=True,
)
return cls(model, tokenizer, collator, device, dtype)
@staticmethod
def _load_trainable(model, state_dict):
"""Load the trainable-only checkpoint, and prove every tensor landed."""
state = {k[len("model."):]: v for k, v in state_dict.items() if k.startswith("model.")}
expected = {n for n, p in model.named_parameters() if p.requires_grad}
missing, unexpected = expected - set(state), set(state) - expected
assert not unexpected, (
f"{len(unexpected)} tensors in the checkpoint match nothing in the model, "
f"e.g. {sorted(unexpected)[:3]} -- the architecture does not match")
assert not missing, (
f"{len(missing)} trainable tensors were not in the checkpoint, "
f"e.g. {sorted(missing)[:3]} -- they would stay randomly initialised")
result = model.load_state_dict(state, strict=False)
assert not result.unexpected_keys, result.unexpected_keys
logger.info(" restored %d trainable tensors (%.1fM params)", len(state),
sum(v.numel() for v in state.values()) / 1e6)
# -- generation --
def _to_rows(self, conversations):
"""DeSTA3-style messages -> the row dicts `Collator` consumes."""
rows = []
for conv in conversations:
audios = []
for message in conv:
for audio in message.get("audios", []):
path = audio["audio"] if isinstance(audio, dict) else audio
assert os.path.exists(path), f"no such audio: {path}"
audios.append({"audio_filepath": path})
n_locators = sum(m["content"].count(self.collator.audio_locator) for m in conv)
assert n_locators == len(audios), (
f"{len(audios)} audios but {n_locators} {self.collator.audio_locator} "
"in the conversation")
# strip `audios` before the chat template ever sees it
rows.append({"messages": [{"role": m["role"], "content": m["content"]} for m in conv],
"audios": audios, "target": ""})
return rows
@torch.no_grad()
def generate(self, conversations, max_new_tokens=200, do_sample=False, **generation_kwargs):
if conversations and isinstance(conversations[0], dict):
conversations = [conversations] # a single conversation
batch = self.collator(self._to_rows(conversations))
batch["input_ids"] = batch["input_ids"].to(self.device)
batch["attention_mask"] = batch["attention_mask"].to(self.device)
# SpeechLLM.encode_audio casts features to the encoder's own dtype
batch["input_features"] = batch["input_features"].to(self.device)
# The trainable tensors (connector, LoRA) stay in float32 while the frozen encoder and
# LLM run in `dtype`, exactly as in training -- so the forward pass has to happen under
# autocast, or layer_norm sees a float32 weight and a float16 activation and raises
# "expected scalar type Half but found Float".
# passing inputs_embeds means `generate` returns only the new tokens --
# there is no prompt prefix to slice off
with torch.autocast(self.device.split(":")[0], dtype=self.dtype,
enabled=not self.device.startswith("cpu")):
generated = self.model.generate(
batch, self.tokenizer, max_new_tokens=max_new_tokens,
do_sample=do_sample, **generation_kwargs)
return [t.strip() for t in
self.tokenizer.batch_decode(generated, skip_special_tokens=True)]
# -- main --
def main():
ap = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--ckpt", required=True)
ap.add_argument("--audio", nargs="+", help="one or more audio files")
ap.add_argument("--manifest", help="jsonl to run over instead of --audio")
ap.add_argument("--audio-root", default="")
ap.add_argument("--limit", type=int, default=10, help="rows to take from --manifest")
ap.add_argument("--prompt", default=DEFAULT_PROMPT)
ap.add_argument("--batch-size", type=int, default=4)
ap.add_argument("--max-new-tokens", type=int, default=200)
ap.add_argument("--device", default=None)
args = ap.parse_args()
assert args.audio or args.manifest, "give --audio or --manifest"
pipe = SpeechLLMForInference.from_checkpoint(args.ckpt, device=args.device)
if args.audio:
items = [(os.path.join(args.audio_root, p), None) for p in args.audio]
else:
items = []
with open(args.manifest) as f:
for line in f:
if len(items) >= args.limit:
break
row = json.loads(line)
items.append((os.path.join(args.audio_root,
row["audios"][0]["audio_filepath"]),
row.get("response") or row.get("target")))
for start in range(0, len(items), args.batch_size):
chunk = items[start:start + args.batch_size]
convs = [[{"role": "user", "content": args.prompt,
"audios": [{"audio": path}]}] for path, _ in chunk]
for (path, reference), hypothesis in zip(chunk, pipe.generate(
convs, max_new_tokens=args.max_new_tokens)):
print(f"\n--- {os.path.basename(path)}")
if reference is not None:
print(f" ref: {reference}")
print(f" hyp: {hypothesis}")
if __name__ == "__main__":
main()
|