File size: 12,114 Bytes
59a1af8 | 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 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 | # Copyright The Marin Authors
# SPDX-License-Identifier: Apache-2.0
"""HF modeling for the SODA hierarchical (backbone + depth) audio LM.
This file is copied verbatim into every exported checkpoint directory and
loaded via trust_remote_code, so it may import only torch/transformers.
The model consumes and produces the same flat frame-interleaved token stream
as the flattened arm (8 audio ids per Mimi frame: semantic then 7 acoustics),
so it is a drop-in for likelihood evals that gather next-token log-probs from
``model(ids).logits`` and for ``model.generate``:
- Internally, positions are grouped into backbone "steps": one text/special
token, or one whole frame (its 8 codebook embeddings summed).
- ``logits[t]`` is the model's true factorized conditional for token ``t+1``:
the 130,308-way unified head (text/special/semantic ids — identical to flat
ids 0..130307) when ``t+1`` starts a step, or the 2,048-way depth head for
codebook k mapped into its flat id block when ``t+1`` is acoustic. All other
vocabulary entries are -inf, so ``log_softmax`` reproduces the factorized
log-prob exactly.
- The depth factorization matches training exactly: codebook k is predicted
from the backbone hidden of the PREVIOUS step plus codebooks 0..k-2 of the
same frame (the immediately preceding codebook is not in the conditioning
set for k >= 2 — a property of the trained shifted-prefix scheme, replicated
verbatim).
"""
from typing import ClassVar
import torch
import torch.nn as nn
from transformers import PreTrainedModel
from transformers.cache_utils import DynamicCache
from transformers.modeling_outputs import CausalLMOutputWithPast
from transformers.models.qwen3.modeling_qwen3 import Qwen3Model
from .configuration_soda_hier import SodaHierConfig
def _group_steps(ids: torch.Tensor, audio_id_lo: int, num_codebooks: int):
"""Map a flat interleaved stream to steps.
Returns (steps, step_of, slot_of): ``steps`` is (S, num_codebooks) with -1
padding on non-frame steps; ``step_of[t]``/``slot_of[t]`` locate position t.
"""
T = ids.shape[0]
step_of = torch.empty(T, dtype=torch.long)
slot_of = torch.empty(T, dtype=torch.long)
rows: list[list[int]] = []
run = 0
for t in range(T):
tok = int(ids[t])
if tok < audio_id_lo:
rows.append([tok] + [-1] * (num_codebooks - 1))
run = 0
else:
if run % num_codebooks == 0:
rows.append([-1] * num_codebooks)
rows[-1][run % num_codebooks] = tok
slot_of[t] = run % num_codebooks
step_of[t] = len(rows) - 1
run += 1
continue
step_of[t] = len(rows) - 1
slot_of[t] = 0
steps = torch.tensor(rows, dtype=torch.long)
return steps, step_of, slot_of
class SodaHierForCausalLM(PreTrainedModel):
config_class = SodaHierConfig
_no_split_modules: ClassVar[list[str]] = ["Qwen3DecoderLayer"]
main_input_name = "input_ids"
_tied_weights_keys: ClassVar[list[str]] = []
def __init__(self, config: SodaHierConfig):
super().__init__(config)
self.backbone = Qwen3Model(config.backbone_config())
self.depth = Qwen3Model(config.depth_config())
e, e_d = config.hidden_size, config.depth_hidden_size
self.unified_head = nn.Linear(e, config.unified_vocab_size, bias=False)
self.bd_proj = nn.Linear(e, e_d, bias=False)
self.acoustic_heads = nn.ModuleList(
nn.Linear(e_d, config.codebook_size, bias=False) for _ in range(config.num_codebooks - 1)
)
self.post_init()
# ------------------------------------------------------------------ core
def _embed_steps(self, steps: torch.Tensor) -> torch.Tensor:
"""(S, num_codebooks) step ids (-1 = empty slot) -> (S, E) summed embeddings."""
valid = steps >= 0
emb = self.backbone.embed_tokens(steps.clamp(min=0))
return (emb * valid.unsqueeze(-1)).sum(dim=1)
def _depth_hidden_for_frames(self, cond: torch.Tensor, frames: torch.Tensor) -> torch.Tensor:
"""Teacher-forced depth pass. cond (F, E_d); frames (F, 8) LM ids -> (F, 8, E_d)."""
cfg = self.config
audio_idx = (frames - cfg.audio_id_lo).clamp(0, cfg.num_codebooks * cfg.codebook_size - 1)
prefix = self.depth.embed_tokens(audio_idx) # (F, 8, E_d)
shifted = torch.roll(prefix, 1, dims=1)
shifted[:, 0] = 0.0
x = cond.unsqueeze(1) + shifted
return self.depth(inputs_embeds=x).last_hidden_state
def _forward_one(self, ids: torch.Tensor) -> torch.Tensor:
"""One unpadded row (T,) -> logits (T, vocab)."""
cfg = self.config
dev = ids.device
steps, step_of, slot_of = _group_steps(ids.cpu(), cfg.audio_id_lo, cfg.num_codebooks)
steps, step_of, slot_of = steps.to(dev), step_of.to(dev), slot_of.to(dev)
T = ids.shape[0]
emb = self._embed_steps(steps) # (S, E)
h = self.backbone(inputs_embeds=emb.unsqueeze(0)).last_hidden_state[0] # (S, E)
u = self.unified_head(h) # (S, unified)
is_audio = ids >= cfg.audio_id_lo
is_frame_step = steps[:, 1] >= 0 # frame steps have slot-1 filled
frame_steps = torch.nonzero(is_frame_step, as_tuple=False)[:, 0]
# depth conditions on the hidden of the step BEFORE the frame
cond_frames = frame_steps[frame_steps >= 1]
d = None
frame_row = torch.full((steps.shape[0],), -1, dtype=torch.long, device=dev)
if len(cond_frames):
d = self._depth_hidden_for_frames(self.bd_proj(h[cond_frames - 1]), steps[cond_frames])
frame_row[cond_frames] = torch.arange(len(cond_frames), device=dev)
logits = torch.full((T, cfg.vocab_size), float("-inf"), dtype=h.dtype, device=dev)
# positions whose NEXT token starts a step: non-audio positions and slot-7 audio positions
primary = (~is_audio) | (slot_of == cfg.num_codebooks - 1)
logits[primary, : cfg.unified_vocab_size] = u[step_of[primary]]
# positions whose next token is acoustic codebook k+1 of the SAME frame
for k in range(cfg.num_codebooks - 1):
pos = torch.nonzero(is_audio & (slot_of == k), as_tuple=False)[:, 0]
if not len(pos):
continue
rows = frame_row[step_of[pos]]
ok = rows >= 0
lo = cfg.audio_id_lo + (k + 1) * cfg.codebook_size
if ok.any():
logits[pos[ok], lo : lo + cfg.codebook_size] = self.acoustic_heads[k](d[rows[ok], k])
if (~ok).any():
# frame at step 0: the factorization defines no conditional; keep finite
logits[pos[~ok]] = 0.0
return logits
def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs) -> CausalLMOutputWithPast:
if input_ids is None:
raise ValueError("SodaHierForCausalLM.forward requires input_ids")
rows = []
for b in range(input_ids.shape[0]):
ids = input_ids[b]
if attention_mask is not None:
length = int(attention_mask[b].sum())
logits_b = torch.zeros((ids.shape[0], self.config.vocab_size), dtype=torch.float32, device=ids.device)
logits_b[:length] = self._forward_one(ids[:length])
rows.append(logits_b)
else:
rows.append(self._forward_one(ids))
logits = torch.stack(rows)
loss = None
if labels is not None:
shift_logits = logits[:, :-1].reshape(-1, self.config.vocab_size)
shift_labels = labels[:, 1:].reshape(-1)
loss = nn.functional.cross_entropy(shift_logits, shift_labels, ignore_index=-100)
return CausalLMOutputWithPast(loss=loss, logits=logits)
# -------------------------------------------------------------- sampling
@staticmethod
def _sample(logits: torch.Tensor, do_sample: bool, temperature: float, top_p: float) -> int:
if not do_sample or temperature <= 0:
return int(logits.argmax())
probs = torch.softmax(logits / temperature, dim=-1)
if top_p is not None and top_p < 1.0:
sorted_probs, sorted_idx = probs.sort(descending=True)
cum = sorted_probs.cumsum(-1)
# drop tokens entirely beyond the nucleus; the crossing token stays
remove = (cum - sorted_probs) > top_p
sorted_probs[remove] = 0.0
sorted_probs /= sorted_probs.sum()
return int(sorted_idx[torch.multinomial(sorted_probs, 1)])
return int(torch.multinomial(probs, 1))
@torch.no_grad()
def generate(
self,
input_ids=None,
attention_mask=None,
max_new_tokens: int = 200,
do_sample: bool = False,
temperature: float = 1.0,
top_p: float = 1.0,
eos_token_id=None,
pad_token_id=None,
**ignored,
) -> torch.Tensor:
"""Two-stage autoregressive decode over the flat interleaved stream.
The backbone advances once per step with a KV cache; whenever the
unified head emits a semantic token, the depth transformer fills in
the frame's 7 acoustic codebooks before the backbone moves on. Only
whole frames are emitted: if fewer than 8 tokens of budget remain
when a frame starts, generation stops early instead.
"""
cfg = self.config
if input_ids.shape[0] != 1:
raise NotImplementedError("SodaHierForCausalLM.generate supports batch size 1")
dev = input_ids.device
eos = set()
if eos_token_id is not None:
eos = {eos_token_id} if isinstance(eos_token_id, int) else set(eos_token_id)
ids = input_ids[0]
steps, _, _ = _group_steps(ids.cpu(), cfg.audio_id_lo, cfg.num_codebooks)
steps = steps.to(dev)
emb = self._embed_steps(steps)
cache = DynamicCache()
out = self.backbone(inputs_embeds=emb.unsqueeze(0), past_key_values=cache, use_cache=True)
cache = out.past_key_values
h_last = out.last_hidden_state[0, -1]
step_pos = steps.shape[0]
generated: list[int] = []
while len(generated) < max_new_tokens:
primary = self._sample(self.unified_head(h_last), do_sample, temperature, top_p)
if primary < cfg.audio_id_lo: # text or special: a one-token step
generated.append(primary)
next_emb = self.backbone.embed_tokens(torch.tensor([primary], device=dev))[0]
if primary in eos:
break
else: # semantic token: emit a whole frame via the depth transformer
if max_new_tokens - len(generated) < cfg.num_codebooks:
break
cond = self.bd_proj(h_last)
frame = [primary]
xs = [cond]
for j in range(cfg.num_codebooks - 1):
d = self.depth(inputs_embeds=torch.stack(xs).unsqueeze(0)).last_hidden_state[0, -1]
idx = self._sample(self.acoustic_heads[j](d), do_sample, temperature, top_p)
frame.append(cfg.audio_id_lo + (j + 1) * cfg.codebook_size + idx)
xs.append(cond + self.depth.embed_tokens(torch.tensor(frame[-1] - cfg.audio_id_lo, device=dev)))
generated.extend(frame)
frame_t = torch.tensor(frame, device=dev)
next_emb = self.backbone.embed_tokens(frame_t).sum(dim=0)
if eos & set(frame):
break
out = self.backbone(
inputs_embeds=next_emb.view(1, 1, -1),
past_key_values=cache,
use_cache=True,
position_ids=torch.tensor([[step_pos]], device=dev),
)
cache = out.past_key_values
h_last = out.last_hidden_state[0, -1]
step_pos += 1
return torch.cat([ids, torch.tensor(generated, dtype=ids.dtype, device=dev)]).unsqueeze(0)
|