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