Views
No views yet
| Models | 🤗 Hugging Face |
|---|---|
| Llama-Mimi-1.3B | llm-jp/Llama-Mimi-1.3B |
| Llama-Mimi-8B | llm-jp/Llama-Mimi-8B |
uv add transformers torch torchaudio1from transformers import AutoModelForCausalLM, AutoTokenizer, MimiModel, AutoFeatureExtractor, StoppingCriteria
2import torch
3import torchaudio
4import re
5import requests
6import io
7
8def audio_array_to_text(
9 audio_array: torch.tensor,
10 audio_tokenizer,
11 feature_extractor,
12 num_quantizers: int,
13) -> str:
14 inputs = feature_extractor(
15 raw_audio=audio_array,
16 sampling_rate=feature_extractor.sampling_rate,
17 return_tensors="pt",
18 ).to(audio_tokenizer.device)
19 with torch.no_grad():
20 encoder_outputs = audio_tokenizer.encode(
21 inputs["input_values"],
22 inputs["padding_mask"],
23 num_quantizers=num_quantizers,
24 )
25 flatten_audio_codes = encoder_outputs.audio_codes.transpose(1, 2).reshape(-1)
26 assert flatten_audio_codes.numel() % num_quantizers == 0
27 steps = []
28 for i in range(0, flatten_audio_codes.numel(), num_quantizers):
29 group = [
30 f"<{flatten_audio_codes[i + j].item()}_{j}>" for j in range(num_quantizers)
31 ]
32 steps.append(group)
33
34 parts = [tok for step in steps for tok in step]
35
36 text = "".join(parts)
37
38 return f"<audio>{text}</audio>"
39
40def text_to_audio_values(
41 text: str,
42 num_quantizers: int,
43 output_file: str,
44 audio_tokenizer,
45 feature_extractor,
46):
47 # Extract (val, idx) pairs from the <val_idx> format in the text
48 matches = re.findall(r"<(\d+)_(\d+)>", text)
49 vals = []
50 for i in range(0, len(matches), num_quantizers):
51 chunk = matches[i : i + num_quantizers]
52 if len(chunk) < num_quantizers:
53 break
54 indices = [int(idx) for _, idx in chunk]
55 if indices == list(range(num_quantizers)):
56 vals.extend(int(val) for val, _ in chunk)
57 else:
58 break
59 vals = vals[: len(vals) - len(vals) % num_quantizers]
60 tensor_bt4 = torch.tensor(vals).reshape(1, -1, num_quantizers) # (B, T, 4)
61 tensor_b4t = tensor_bt4.transpose(1, 2) # (B, 4, T)
62 audio_values = audio_tokenizer.decode(tensor_b4t)[0]
63 torchaudio.save(
64 output_file,
65 audio_values[0].detach().cpu(),
66 feature_extractor.sampling_rate,
67 )
68
69
70class StopOnAudioEnd(StoppingCriteria):
71 def __init__(self, tokenizer):
72 self.tokenizer = tokenizer
73 self.target_text = "</audio>"
74 self.target_ids = tokenizer(
75 self.target_text, add_special_tokens=False
76 ).input_ids
77
78 def __call__(self, input_ids, scores, **kwargs):
79 if len(input_ids[0]) < len(self.target_ids):
80 return False
81 return input_ids[0][-len(self.target_ids) :].tolist() == self.target_ids
82
83temperature = 0.8
84top_k = 30
85do_sample = True
86max_length = 1024
87device = "cuda" if torch.cuda.is_available() else "cpu"
88model_id = "llm-jp/Llama-Mimi-1.3B"
89model = AutoModelForCausalLM.from_pretrained(model_id, torch_dtype=torch.bfloat16).eval().to(device)
90num_quantizers = model.config.num_quantizers
91tokenizer = AutoTokenizer.from_pretrained(model_id)
92audio_tokenizer = MimiModel.from_pretrained("kyutai/mimi")
93feature_extractor = AutoFeatureExtractor.from_pretrained("kyutai/mimi")
94stopping_criteria = StopOnAudioEnd(tokenizer)
95
96audio_url = "https://speed1313.github.io/llama-mimi/data/prompt/natural/great_day_gt.wav"
97response = requests.get(audio_url)
98response.raise_for_status()
99waveform, sample_rate = torchaudio.load(io.BytesIO(response.content))
100if sample_rate != feature_extractor.sampling_rate:
101 waveform = torchaudio.transforms.Resample(
102 sample_rate, feature_extractor.sampling_rate
103 )(waveform)
104 sample_rate = feature_extractor.sampling_rate
105prompt_array = waveform.squeeze().cpu().numpy()
106
107text = audio_array_to_text(
108 prompt_array, audio_tokenizer, feature_extractor, num_quantizers
109)
110
111text = text.replace("</audio>", "")
112inputs = tokenizer(text, return_tensors="pt").to(device)
113
114with torch.no_grad():
115 generated = model.generate(
116 **inputs,
117 max_length=max_length,
118 do_sample=do_sample,
119 temperature=temperature,
120 top_k=top_k,
121 stopping_criteria=[stopping_criteria],
122 )
123
124generated_text = tokenizer.decode(generated[0])
125
126text_to_audio_values(
127 generated_text,
128 num_quantizers=num_quantizers,
129 output_file="output.wav",
130 audio_tokenizer=audio_tokenizer,
131 feature_extractor=feature_extractor,
132)
133@misc{sugiura2025llamamimispeechlanguagemodels,
title={Llama-Mimi: Speech Language Models with Interleaved Semantic and Acoustic Tokens},
author={Issa Sugiura and Shuhei Kurita and Yusuke Oda and Ryuichiro Higashinaka},
year={2025},
eprint={2509.14882},
archivePrefix={arXiv},
primaryClass={cs.CL},
url={https://arxiv.org/abs/2509.14882},
}