1use crate::config::GenerateConfig;
2use crate::lm::LmModel;
3use crate::sampling::LogitsProcessor;
4use anyhow::{Result, ensure};
5
6pub const UNGENERATED: u32 = u32::MAX;
7
8pub struct GenerateState {
10 audio_tokens: Vec<Vec<u32>>,
11 text_tokens: Vec<u32>,
12 text_lp: LogitsProcessor,
13 audio_lp: LogitsProcessor,
14 step_idx: usize,
15 forced_audio_tokens: ForcedAudioTokens,
16 cfg: GenerateConfig,
17}
18
19#[derive(Debug, Clone)]
20pub(crate) struct ForcedAudioTokens {
21 delay: usize,
22 pad: u32,
23 pattern: Vec<usize>,
24}
25
26impl ForcedAudioTokens {
27 pub(crate) fn new(delay: usize, pad: u32, pattern: &[usize]) -> Self {
28 Self {
29 delay,
30 pad,
31 pattern: pattern.to_vec(),
32 }
33 }
34
35 pub(crate) fn forced_tokens(&self, step: usize) -> Vec<Option<u32>> {
36 if step >= self.delay {
37 return vec![None; self.pattern.len()];
38 }
39 self.pattern
40 .iter()
41 .map(|&v| if v == 0 { None } else { Some(self.pad) })
42 .collect()
43 }
44}
45
46impl GenerateState {
47 pub fn new(
48 max_steps: usize,
49 text_lp: LogitsProcessor,
50 audio_lp: LogitsProcessor,
51 cfg: GenerateConfig,
52 ) -> Self {
53 let buf = max_steps + cfg.acoustic_delay;
54 let audio_tokens = vec![vec![UNGENERATED; cfg.total_audio_codebooks()]; buf];
55 let text_tokens = vec![UNGENERATED; buf];
56 let forced = ForcedAudioTokens::new(cfg.acoustic_delay, cfg.audio_pad_token(), &[8, 8]);
57 Self {
58 audio_tokens,
59 text_tokens,
60 text_lp,
61 audio_lp,
62 step_idx: 0,
63 forced_audio_tokens: forced,
64 cfg,
65 }
66 }
67
68 pub fn config(&self) -> &GenerateConfig {
69 &self.cfg
70 }
71
72 pub fn step_idx(&self) -> usize {
73 self.step_idx
74 }
75
76 pub fn text_tokens(&self) -> &[u32] {
77 let n = self.step_idx.min(self.text_tokens.len());
78 &self.text_tokens[..n]
79 }
80
81 pub fn step(&mut self, lm: &mut LmModel, text_token: u32, input_audio: &[u32]) -> Result<u32> {
83 ensure!(
84 input_audio.len() == self.cfg.input_audio_codebooks,
85 "expected {} user codebooks, got {}",
86 self.cfg.input_audio_codebooks,
87 input_audio.len()
88 );
89 for (ci, &t) in input_audio.iter().enumerate() {
90 let idx = ci + self.cfg.generated_audio_codebooks;
91 self.audio_tokens[self.step_idx][idx] = t;
92 }
93 let pad = self.cfg.audio_pad_token();
94 let mut delayed = Vec::with_capacity(self.cfg.total_audio_codebooks());
95 for codebook in 0..self.cfg.total_audio_codebooks() {
96 let t = if codebook == 0 || codebook == self.cfg.generated_audio_codebooks {
97 if self.step_idx == 0 {
98 pad
99 } else {
100 self.audio_tokens[self.step_idx - 1][codebook]
101 }
102 } else if self.step_idx <= self.cfg.acoustic_delay {
103 pad
104 } else {
105 self.audio_tokens[self.step_idx - self.cfg.acoustic_delay - 1][codebook]
106 };
107 ensure!(
108 t != UNGENERATED,
109 "internal: ungenerated audio at step {}",
110 self.step_idx
111 );
112 delayed.push(Some(t));
113 }
114 let (text_logits, hidden) = lm.forward_step(Some(text_token), &delayed)?;
115 let sampled_text = self.text_lp.sample(text_logits.view())?;
116 self.text_tokens[self.step_idx] = sampled_text;
117 let forced = self.forced_audio_tokens.forced_tokens(self.step_idx);
118 if let Some(tokens) =
119 lm.depformer_sample(&hidden, Some(sampled_text), &forced, &mut self.audio_lp)?
120 {
121 for (ci, &tok) in tokens.iter().enumerate() {
122 let delay = if ci == 0 { 0 } else { self.cfg.acoustic_delay };
123 let pos = self.step_idx.saturating_sub(delay);
124 self.audio_tokens[pos][ci] = tok;
125 }
126 }
127 self.step_idx += 1;
128 Ok(sampled_text)
129 }
130
131 pub fn last_audio_frame(&self) -> Option<Vec<u32>> {
133 if self.step_idx <= self.cfg.acoustic_delay {
134 return None;
135 }
136 let pos = self.step_idx - self.cfg.acoustic_delay - 1;
137 let frame = &self.audio_tokens[pos];
138 let pad = self.cfg.audio_pad_token();
139 if frame[..self.cfg.generated_audio_codebooks]
140 .iter()
141 .any(|&t| t >= pad)
142 {
143 return None;
144 }
145 Some(frame[..self.cfg.generated_audio_codebooks].to_vec())
146 }
147
148 pub fn reset(&mut self, lm: &mut LmModel) {
149 lm.reset_state();
150 self.step_idx = 0;
151 let buf = self.audio_tokens.len();
152 let tc = self.cfg.total_audio_codebooks();
153 self.audio_tokens = vec![vec![UNGENERATED; tc]; buf];
154 self.text_tokens = vec![UNGENERATED; self.text_tokens.len()];
155 }
156}