File size: 3,346 Bytes
724479b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Standalone inference for Pranavz/583."""
import argparse
import re
import wave
from pathlib import Path

import torch
from huggingface_hub import snapshot_download
from transformers import AutoFeatureExtractor, AutoModelForCausalLM, AutoTokenizer, MimiModel

AUDIO_RE = re.compile(r"^<(\d+)_(\d+)>$")


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--repo-id", default="Pranavz/583")
    p.add_argument("--speaker", choices=["Ira", "Aisha", "Siya", "Zoya", "Silver"], default="Ira")
    p.add_argument("--text", required=True)
    p.add_argument("--output", default="output.wav")
    p.add_argument("--temperature", type=float, default=0.8)
    p.add_argument("--top-k", type=int, default=30)
    p.add_argument("--max-new-tokens", type=int, default=2048)
    p.add_argument("--device", default="cuda")
    args = p.parse_args()

    root = Path(snapshot_download(args.repo_id))
    dtype = torch.bfloat16 if args.device.startswith("cuda") else torch.float32
    tokenizer = AutoTokenizer.from_pretrained(root)
    model = AutoModelForCausalLM.from_pretrained(
        root, trust_remote_code=True, torch_dtype=dtype, attn_implementation="sdpa"
    ).eval().to(args.device)
    mimi = MimiModel.from_pretrained(root / "codec", torch_dtype=dtype).eval().to(args.device)
    feature = AutoFeatureExtractor.from_pretrained(root / "codec")

    vocab = tokenizer.get_vocab()
    mapping = {}
    for token, token_id in vocab.items():
        match = AUDIO_RE.match(token)
        if match:
            mapping[int(token_id)] = (int(match.group(1)), int(match.group(2)))
    allowed = torch.tensor(sorted([*mapping, int(vocab["</audio>"])]), device=args.device)
    prompt = f'<text>{args.speaker}: {args.text}<audio>'
    inputs = tokenizer(prompt, return_tensors="pt").to(args.device)
    output = model.generate_audio(
        **inputs,
        allowed_ids=allowed,
        max_new_tokens=args.max_new_tokens,
        min_new_tokens=8,
        temperature=args.temperature,
        top_k=args.top_k,
        do_sample=True,
    )[0].tolist()

    start = len(inputs.input_ids[0])
    try:
        end = output.index(int(vocab["</audio>"]), start)
    except ValueError:
        end = len(output)
    values, frame, expected = [], [], 0
    for token_id in output[start:end]:
        item = mapping.get(int(token_id))
        if item is None:
            frame, expected = [], 0
            continue
        code, q = item
        if q == expected:
            frame.append(code)
            expected += 1
            if expected == 8:
                values.extend(frame)
                frame, expected = [], 0
        elif q == 0:
            frame, expected = [code], 1
        else:
            frame, expected = [], 0
    if not values:
        raise RuntimeError("No complete Mimi-Q8 frames generated")
    codes = torch.tensor(values, device=args.device).reshape(1, -1, 8).transpose(1, 2)
    audio = mimi.decode(codes).audio_values[0, 0].float().cpu().clamp(-1, 1)
    pcm = (audio.numpy() * 32767).astype("<i2")
    with wave.open(args.output, "wb") as f:
        f.setnchannels(1); f.setsampwidth(2); f.setframerate(int(feature.sampling_rate)); f.writeframes(pcm.tobytes())
    print(f"Saved {args.output} ({len(audio) / feature.sampling_rate:.2f}s)")


if __name__ == "__main__":
    main()