"""Speaker-conditioned inference for rumik-oss 1 base.""" import argparse import wave from pathlib import Path import torch from huggingface_hub import snapshot_download from transformers import ( AutoFeatureExtractor, AutoModelForCausalLM, AutoTokenizer, MimiModel, ) def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--model", default=str(Path(__file__).resolve().parent)) parser.add_argument("--speaker", default="Ira") parser.add_argument("--text", required=True) parser.add_argument("--output", default="speech.wav") parser.add_argument("--device", default="cuda") parser.add_argument("--temperature", type=float, default=0.8) parser.add_argument("--top-k", type=int, default=30) parser.add_argument("--max-new-tokens", type=int, default=2048) parser.add_argument("--seed", type=int, default=42) args = parser.parse_args() if not args.text.strip(): parser.error("--text must not be empty") if args.temperature <= 0 or args.top_k < 0 or args.max_new_tokens < 8: parser.error( "temperature must be positive, top-k nonnegative, and max-new-tokens at least 8" ) root = Path(args.model) if not root.is_dir(): root = Path(snapshot_download(args.model)) torch.manual_seed(args.seed) dtype = torch.bfloat16 if args.device.startswith("cuda") else torch.float32 tokenizer = AutoTokenizer.from_pretrained(root, trust_remote_code=True) model = ( AutoModelForCausalLM.from_pretrained( root, trust_remote_code=True, dtype=dtype, attn_implementation="sdpa" ) .eval() .to(args.device) ) codec = MimiModel.from_pretrained(root / "codec").eval().to(args.device) sample_rate = AutoFeatureExtractor.from_pretrained(root / "codec").sampling_rate speakers = tuple(model.config.speakers) if args.speaker not in speakers: parser.error(f"--speaker must be one of: {', '.join(speakers)}") inputs = tokenizer( f"{args.speaker}: {args.text}