Views
No views yet
pip install -U torch torchaudio transformerstrust_remote_code=True.1import torch, torchaudio
2from transformers import AutoModel
3
4device = "cuda" if torch.cuda.is_available() else "cpu"
5
6# 1) Load the WavCoch tokenizer (audio -> token IDs)
7quantizer = AutoModel.from_pretrained(
8 "TuKoResearch/WavCochV8192", trust_remote_code=True
9).to(device).eval()
10
11# 2) Load the AuriStream LM (tokens -> hidden states / next-token prediction)
12lm = AutoModel.from_pretrained(
13 "TuKoResearch/AuriStream1B_librilight_ckpt500k", trust_remote_code=True
14).to(device).eval()
15
16# 3) Read an audio file (mono, 16 kHz recommended)
17wav, sr = torchaudio.load("sample.wav")
18
19if wav.size(0) > 1: # stereo -> mono
20 wav = wav.mean(dim=0, keepdim=True)
21if sr != 16_000:
22 wav = torchaudio.transforms.Resample(sr, 16_000)(wav)
23 sr = 16_000
24
25# 4) Quantize the audio to obtain cochlear token IDs
26with torch.no_grad():
27 # The quantizer forward method expects (B, 1, T); returns (B, L)
28 token_ids = quantizer(wav.unsqueeze(0).to(device))['input_ids'] # (1, L)
29
30# 5) Forward pass to obtain hidden states
31with torch.no_grad():
32 out = lm(token_ids, output_hidden_states=True)
33 last_layer = out["hidden_states"][-1] # (1, T, D)
34 last_layer_mean = last_layer.mean(dim=1) # time mean-pool -> (1, D)
35
36print("Mean-pooled embedding shape:", last_layer_mean.shape)output_hidden_states=True returns all layers.1import torch, torchaudio
2from transformers import AutoModel
3
4device = "cuda" if torch.cuda.is_available() else "cpu"
5
6# WavCoch tokenizer (audio -> tokens)
7quantizer = AutoModel.from_pretrained(
8 "TuKoResearch/WavCochV8192", trust_remote_code=True
9).to(device).eval()
10
11# AuriStream LM (tokens -> next tokens)
12lm = AutoModel.from_pretrained(
13 "TuKoResearch/AuriStream1B_librilight_ckpt500k", trust_remote_code=True
14).to(device).eval()
15
16# Load and prep a short prompt (e.g., 3s of audio at 16 kHz)
17prompt_seconds = 3
18wav, sr = torchaudio.load("prompt.wav")
19if wav.size(0) > 1:
20 wav = wav.mean(dim=0, keepdim=True)
21if sr != 16_000:
22 wav = torchaudio.transforms.Resample(sr, 16_000)(wav)
23 sr = 16_000
24# Slice using an integer number of samples
25n_samples = int(round(sr * prompt_seconds))
26wav = wav[:, :n_samples]
27
28# Quantize the prompt audio to get token IDs
29with torch.no_grad():
30 prompt_tokens = quantizer(wav.unsqueeze(0).to(device))['input_ids'] # (1, L)
31
32# Decide how many future tokens to generate ("roll-out")
33tokens_per_sec = prompt_tokens.size(1) / float(prompt_seconds)
34rollout_seconds = 2
35rollout_steps = int(round(tokens_per_sec * rollout_seconds)) # K
36
37# Generate future tokens
38with torch.no_grad():
39 # returns (pred_tokens, pred_logits); temperature/top_k/top_p/seed optional
40 pred_tokens, _ = lm.generate(
41 prompt_tokens, rollout_steps, temp=0.7, top_k=50, top_p=0.95, seed=0
42 )
43 full_tokens = torch.cat([prompt_tokens, pred_tokens], dim=1) # (1, L+K)
1@inproceedings{tuckute2025cochleartokens,
2 title = {Representing Speech Through Autoregressive Prediction of Cochlear Tokens},
3 author = {Greta Tuckute and Klemen Kotar and Evelina Fedorenko and Daniel Yamins},
4 booktitle = {Interspeech 2025},
5 year = {2025},
6 pages = {2180--2184},
7 doi = {10.21437/Interspeech.2025-2044},
8 issn = {2958-1796}
9}