| |
| """ |
| 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 |
|
|
| |
| @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 = 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)) |
|
|
| |
| 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, |
| |
| 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) |
|
|
| |
| 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") |
| |
| 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] |
| 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) |
| |
| batch["input_features"] = batch["input_features"].to(self.device) |
|
|
| |
| |
| |
| |
| |
| |
| 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)] |
|
|
|
|
| |
|
|
| 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() |
|
|