| import os |
| import math |
| import torch |
| import warnings |
| import logging |
| import contextlib |
| import io |
| from torch import nn |
| from torch.nn import functional as F |
| from transformers.modeling_outputs import MoeCausalLMOutputWithPast |
| from transformers import SiglipVisionModel, SiglipImageProcessor, logging as hf_logging |
|
|
| from core import RMSNorm, precompute_freqs_cis, Block, MOEFeedForward |
| from models.lm.config import LMConfig |
| from models.lm.model import LMForCausalLM |
| from models.vam.config import VAMConfig |
| from encoders.audio import SenseVoiceAudioEncoder, SenseVoiceAudioProcessor |
| from encoders.vision import SiglipVisionEncoder |
| from projectors import MMVisionProjector, MMAudioProjector |
|
|
|
|
| class TalkerHead(nn.Module): |
| def __init__(self, in_features, out_features, num_layers=8, rank=256): |
| super().__init__() |
| self.num_layers = num_layers |
| self.base = nn.Linear(in_features, out_features, bias=False) |
| self.adapters = nn.ModuleList([ |
| nn.Sequential(nn.Linear(in_features, rank, bias=False), nn.GELU(), nn.Linear(rank, out_features, bias=False)) |
| for _ in range(num_layers) |
| ]) |
|
|
| def forward(self, x): |
| base_out = self.base(x) |
| return [base_out + adapter(x) for adapter in self.adapters] |
|
|
|
|
| class TalkerEmbedding(nn.Module): |
| def __init__(self, num_embeddings, embedding_dim, num_layers=8, rank=256): |
| super().__init__() |
| self.num_layers = num_layers |
| self.base = nn.Embedding(num_embeddings, embedding_dim) |
| self.adapters = nn.ModuleList([ |
| nn.Sequential(nn.Embedding(num_embeddings, rank), nn.GELU(), nn.Linear(rank, embedding_dim, bias=False)) |
| for _ in range(num_layers) |
| ]) |
|
|
| def forward(self, x): |
| base_out = self.base(x) |
| return sum(base_out[:, i, :] + self.adapters[i](x[:, i, :]) for i in range(len(self.adapters))) / self.num_layers |
|
|
|
|
| class TalkerModule(nn.Module): |
| def __init__(self, config: VAMConfig): |
| super().__init__() |
| self.talker_config = LMConfig(hidden_size=config.talker_hidden_size, use_moe=config.use_moe) |
| self.layers = nn.ModuleList([Block(l, self.talker_config) for l in range(config.num_talker_hidden_layers)]) |
| self.norm = RMSNorm(config.talker_hidden_size, eps=config.rms_norm_eps) |
| self.lm_head = TalkerHead(config.talker_hidden_size, config.audio_vocab_size) |
| self.embed_tokens = TalkerEmbedding(config.audio_vocab_size, config.talker_hidden_size) |
| self.codec_proj = nn.Sequential( |
| nn.Linear(config.talker_hidden_size, config.talker_hidden_size), |
| nn.GELU(), |
| nn.Linear(config.talker_hidden_size, config.talker_hidden_size), |
| RMSNorm(config.talker_hidden_size, eps=config.rms_norm_eps), |
| ) |
| self.embed_proj = nn.Sequential( |
| nn.Linear(config.hidden_size, config.hidden_size), |
| nn.GELU(), |
| nn.Linear(config.hidden_size, config.talker_hidden_size), |
| RMSNorm(config.talker_hidden_size, eps=config.rms_norm_eps), |
| ) |
| self.text_scale, self.audio_scale = nn.Parameter(torch.tensor(3.0)), nn.Parameter(torch.tensor(1.0)) |
| self.spk_proj = nn.Linear(config.spk_emb_size, config.talker_hidden_size, bias=False) |
| freqs_cos, freqs_sin = precompute_freqs_cis( |
| dim=self.talker_config.head_dim, end=config.max_position_embeddings, |
| rope_base=config.rope_theta, rope_scaling=config.rope_scaling |
| ) |
| self.register_buffer("freqs_cos", freqs_cos, persistent=False) |
| self.register_buffer("freqs_sin", freqs_sin, persistent=False) |
|
|
|
|
| class VAM(LMForCausalLM): |
| config_class = VAMConfig |
|
|
| def __init__(self, config: VAMConfig = None, audio_encoder_path: str = None, vision_model_path: str = None): |
| config = config or VAMConfig() |
| super().__init__(config) |
| object.__setattr__(self, 'thinker', self.model) |
| object.__setattr__(self.model, 'lm_head', self.lm_head) |
| self.talker = TalkerModule(config) |
| self.audio_proj = MMAudioProjector(config.audio_hidden_size, config.hidden_size) |
| self.vision_proj = MMVisionProjector(config.image_hidden_size, config.hidden_size, target_tokens=config.image_token_len) |
| self.audio_pad_token, self.audio_stop_token, self.audio_spk_token = config.audio_pad_token, config.audio_stop_token, config.audio_spk_token |
| meta_init = any(p.device.type == 'meta' for p in self.parameters()) |
| if meta_init: |
| object.__setattr__(self, 'audio_encoder', None) |
| object.__setattr__(self, 'audio_processor', None) |
| object.__setattr__(self, 'vision_encoder', None) |
| object.__setattr__(self, 'vision_processor', None) |
| else: |
| audio_enc = SenseVoiceAudioEncoder(audio_encoder_path) if audio_encoder_path else SenseVoiceAudioEncoder() |
| object.__setattr__(self, 'audio_encoder', audio_enc) |
| object.__setattr__(self, 'audio_processor', audio_enc.processor) |
| vision_enc = SiglipVisionEncoder(vision_model_path) if vision_model_path else SiglipVisionEncoder() |
| object.__setattr__(self, 'vision_encoder', vision_enc) |
| object.__setattr__(self, 'vision_processor', vision_enc.processor) |
|
|
| @classmethod |
| def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs): |
| audio_encoder_path = kwargs.pop('audio_encoder_path', None) |
| vision_model_path = kwargs.pop('vision_model_path', None) |
| model = super().from_pretrained(pretrained_model_name_or_path, *model_args, **kwargs) |
| if audio_encoder_path and model.audio_encoder is None: |
| enc = SenseVoiceAudioEncoder(audio_encoder_path) |
| object.__setattr__(model, 'audio_encoder', enc) |
| object.__setattr__(model, 'audio_processor', enc.processor) |
| if vision_model_path and model.vision_encoder is None: |
| vision_enc = SiglipVisionEncoder(vision_model_path) |
| object.__setattr__(model, 'vision_encoder', vision_enc) |
| object.__setattr__(model, 'vision_processor', vision_enc.processor) |
| return model |
|
|
| @staticmethod |
| def load_sensevoice(path): |
| if not os.path.exists(path): |
| warnings.warn(f"[VAM] SenseVoice path not found: {path}") |
| return None, None |
| logging.getLogger().setLevel(logging.ERROR) |
| hf_logging.set_verbosity_error() |
| with contextlib.redirect_stdout(io.StringIO()): |
| from funasr import AutoModel |
| m = AutoModel(model=path, trust_remote_code=True, disable_update=True, device="cpu") |
| encoder, frontend = m.model.encoder, m.kwargs["frontend"] |
| for p in encoder.parameters(): |
| p.requires_grad = False |
| return encoder.eval().float(), SenseVoiceAudioProcessor(frontend.eval()) |
|
|
| @staticmethod |
| def load_vision(path): |
| if path is None or not os.path.exists(path): |
| warnings.warn(f"[VAM] Vision model path not found: {path}. vision_encoder will be None!") |
| return None, None |
| hf_logging.set_verbosity_error() |
| try: |
| model = SiglipVisionModel.from_pretrained(path) |
| except (RuntimeError, ValueError): |
| return None, None |
| processor = SiglipImageProcessor.from_pretrained(path) |
| for p in model.parameters(): |
| p.requires_grad = False |
| return model.eval(), processor |
|
|
| @torch.compiler.disable |
| def encode_audio_inputs(self, audio_inputs, audio_lens=None): |
| if (audio_inputs is None) or (self.audio_encoder is None) or (not audio_inputs.any()): |
| return None |
| batch_mask = audio_inputs.flatten(1).any(1) |
| enc_dtype = next(self.audio_encoder.parameters()).dtype |
| valid_fbank = audio_inputs[batch_mask].to(dtype=enc_dtype) |
| if audio_lens is not None: |
| valid_lens = audio_lens[batch_mask].to(valid_fbank.device) |
| else: |
| valid_lens = torch.tensor([valid_fbank.size(1)] * valid_fbank.size(0), device=valid_fbank.device) |
| with torch.no_grad(): |
| emb, _ = self.audio_encoder.model(valid_fbank, valid_lens) |
| proj_dtype = next(self.audio_proj.parameters()).dtype |
| emb_list = [self.audio_proj(emb[i, :max(1, min(valid_lens[i].item(), emb.size(1)))].unsqueeze(0).to(proj_dtype)).squeeze(0) for i in range(emb.size(0))] |
| if batch_mask.all(): |
| return emb_list |
| out = [None] * audio_inputs.size(0) |
| j = 0 |
| for i in range(audio_inputs.size(0)): |
| if batch_mask[i]: |
| out[i] = emb_list[j] |
| j += 1 |
| return out |
|
|
| @torch.compiler.disable |
| def inject_audio_features(self, tokens, h, audio_feats, seqlen): |
| if audio_feats is None or not self.config.audio_ids: |
| return h |
| marker = self.config.audio_ids[0] |
| out = [] |
| for b in range(h.size(0)): |
| hb, seq, i = h[b], tokens[b].tolist(), 0 |
| af = audio_feats[b] if audio_feats[b] is not None else None |
| while i < len(seq): |
| if seq[i] == marker: |
| start = i |
| while i < len(seq) and seq[i] == marker: |
| i += 1 |
| if af is not None: |
| inject_len = min(af.size(0), i - start) |
| hb = torch.cat((hb[:start], af[:inject_len], hb[start + inject_len:]), dim=0) |
| af = None |
| else: |
| i += 1 |
| out.append(hb) |
| return torch.stack(out) |
|
|
| @torch.compiler.disable |
| def get_image_embeddings(self, image_inputs): |
| if hasattr(image_inputs, 'keys'): |
| image_inputs = {k: (v.squeeze(1) if v.ndim > 2 and v.shape[1] == 1 else v) for k, v in image_inputs.items()} |
| pixel_attention_mask = image_inputs.get('pixel_attention_mask') |
| if pixel_attention_mask is not None and not pixel_attention_mask.any(): |
| pv = image_inputs['pixel_values'] |
| return pv.new_zeros(pv.size(0), pv.size(1), self.config.image_hidden_size) |
| with torch.no_grad(): |
| outputs = self.vision_encoder.model(**image_inputs) |
| return outputs.last_hidden_state |
|
|
| @torch.compiler.disable |
| def encode_image_inputs(self, pixel_values): |
| if pixel_values is None or self.vision_encoder is None: |
| return None |
| mask = pixel_values.flatten(1).any(1) |
| if not mask.any(): |
| return pixel_values.new_zeros(pixel_values.size(0), self.config.image_token_len, self.config.hidden_size) |
| with torch.no_grad(): |
| emb = self.vision_encoder.model(pixel_values=pixel_values[mask]).last_hidden_state |
| if emb.dim() == 2: |
| emb = emb.unsqueeze(0) |
| emb = self.vision_proj(emb) |
| if mask.all(): |
| return emb |
| idx = mask.nonzero().view(-1, 1, 1).expand_as(emb) |
| return emb.new_zeros(pixel_values.size(0), *emb.shape[1:]).scatter(0, idx, emb) |
|
|
| @torch.compiler.disable |
| def count_vision_proj(self, tokens, h, vision_tensors=None, seqlen=512): |
| if vision_tensors is None or not self.config.image_ids: |
| return h |
| marker, vf = self.config.image_ids[0], vision_tensors |
| if vf.dim() == 3: |
| vf = vf.unsqueeze(1) |
| out = [] |
| for b in range(h.size(0)): |
| hb, seq, k, i = h[b], tokens[b].tolist(), 0, 0 |
| while i < len(seq): |
| if seq[i] == marker: |
| start = i |
| while i < len(seq) and seq[i] == marker: |
| i += 1 |
| if k < vf.size(1): |
| hb = torch.cat((hb[:start], vf[b][k][:i - start], hb[i:]), dim=0)[:seqlen] |
| k += 1 |
| else: |
| i += 1 |
| out.append(hb) |
| return torch.stack(out) |
|
|
| def forward(self, input_ids, attention_mask=None, past_key_values=None, use_cache=False, logits_to_keep=0, |
| audio_inputs=None, audio_lens=None, pixel_values=None, **args): |
| if len(input_ids.shape) == 2: |
| batch_size, seq_length = input_ids.shape |
| text_ids = input_ids |
| audio_ids = torch.full((batch_size, 8, seq_length), self.audio_pad_token, dtype=torch.long, device=input_ids.device) |
| else: |
| batch_size, _, seq_length = input_ids.shape |
| text_ids, audio_ids = input_ids[:, 8, :], input_ids[:, :8, :] |
| if hasattr(past_key_values, 'layers'): |
| past_key_values = None |
| n_thinker, n_talker = len(self.thinker.layers), len(self.talker.layers) |
| past_key_values = past_key_values or ([None] * (n_thinker + n_talker)) |
| start_pos = past_key_values[0][0].shape[1] if past_key_values[0] is not None else 0 |
| if self.thinker.freqs_cos[0, 0] == 0: |
| freqs_cos, freqs_sin = precompute_freqs_cis(dim=self.config.head_dim, end=self.config.max_position_embeddings, rope_base=self.config.rope_theta, rope_scaling=self.config.rope_scaling) |
| self.thinker.freqs_cos, self.thinker.freqs_sin = freqs_cos.to(input_ids.device), freqs_sin.to(input_ids.device) |
| if self.talker.freqs_cos[0, 0] == 0: |
| freqs_cos, freqs_sin = precompute_freqs_cis(dim=self.talker.talker_config.head_dim, end=self.config.max_position_embeddings, rope_base=self.config.rope_theta, rope_scaling=self.config.rope_scaling) |
| self.talker.freqs_cos, self.talker.freqs_sin = freqs_cos.to(input_ids.device), freqs_sin.to(input_ids.device) |
| presents = [] |
|
|
| hidden_states = self.thinker.dropout(self.thinker.embed_tokens(text_ids)) |
| position_embeddings = (self.thinker.freqs_cos[start_pos:start_pos + seq_length], self.thinker.freqs_sin[start_pos:start_pos + seq_length]) |
| if audio_inputs is not None and start_pos == 0: |
| audio_features = self.encode_audio_inputs(audio_inputs, audio_lens) |
| hidden_states = self.inject_audio_features(text_ids, hidden_states, audio_features, seq_length) |
| if pixel_values is not None and start_pos == 0: |
| if hasattr(pixel_values, 'keys'): |
| img_emb = self.get_image_embeddings(pixel_values).to(hidden_states.dtype) |
| vision_tensors = self.vision_proj(img_emb) |
| else: |
| if len(pixel_values.shape) == 6: |
| pixel_values = pixel_values.squeeze(2) |
| if len(pixel_values.shape) == 4: |
| pixel_values = pixel_values.unsqueeze(1) |
| bs, num, c, im_h, im_w = pixel_values.shape |
| stack_dim = 1 if bs > 1 else 0 |
| vision_tensors = torch.stack([self.encode_image_inputs(pixel_values[:, i, :, :, :]) for i in range(num)], dim=stack_dim) |
| hidden_states = self.count_vision_proj(tokens=text_ids, h=hidden_states, vision_tensors=vision_tensors, seqlen=seq_length) |
| bridge_states = hidden_states |
| for i, (layer, past_key_value) in enumerate(zip(self.thinker.layers, past_key_values[:n_thinker])): |
| hidden_states, present = layer(hidden_states, position_embeddings, past_key_value=past_key_value, use_cache=use_cache, attention_mask=attention_mask) |
| presents.append(present) |
| if i == self.config.bridge_layer: |
| bridge_states = hidden_states |
| h_thinker = self.thinker.norm(hidden_states) |
|
|
| talker_emb = self.talker.embed_tokens(audio_ids) |
| spk_emb = args.get('spk_emb', None) |
| if spk_emb is not None: |
| spk_mask = (audio_ids[:, 0, :] == self.audio_spk_token).unsqueeze(-1) |
| talker_emb = torch.where(spk_mask, self.talker.spk_proj(spk_emb).unsqueeze(1), talker_emb) |
| hidden_states = self.talker.embed_proj(bridge_states) * self.talker.text_scale + self.talker.codec_proj(talker_emb) * self.talker.audio_scale |
| talker_pos_emb = (self.talker.freqs_cos[start_pos:start_pos + seq_length], self.talker.freqs_sin[start_pos:start_pos + seq_length]) |
| for layer, past_key_value in zip(self.talker.layers, past_key_values[n_thinker:]): |
| hidden_states, present = layer(hidden_states, talker_pos_emb, past_key_value=past_key_value, use_cache=use_cache, attention_mask=attention_mask) |
| presents.append(present) |
| h_talker = self.talker.norm(hidden_states) |
|
|
| slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep |
| aux_loss = sum(l.mlp.aux_loss for l in list(self.thinker.layers) + list(self.talker.layers) if isinstance(l.mlp, MOEFeedForward)) |
| aux_loss += sum(p.sum() for p in self.audio_proj.parameters()) * 0 + sum(p.sum() for p in self.vision_proj.parameters()) * 0 + sum(p.sum() for p in self.talker.lm_head.adapters.parameters()) * 0 + sum(p.sum() for p in self.talker.spk_proj.parameters()) * 0 |
| text_logits = self.thinker.lm_head(h_thinker[:, slice_indices, :]) |
| audio_logits = self.talker.lm_head(h_talker[:, slice_indices, :]) |
|
|
| out = MoeCausalLMOutputWithPast(aux_loss=aux_loss, logits=text_logits, past_key_values=presents) |
| out.audio_logits = audio_logits |
| return out |
|
|
| @torch.inference_mode() |
| def generate(self, input_ids, eos_token_id=2, max_new_tokens=1024, temperature=0.75, top_p=0.90, |
| stream=False, rp=1., use_cache=True, return_audio_codes=False, **args): |
| if stream: |
| return self.stream_generate(input_ids, eos_token_id, max_new_tokens, temperature, top_p, rp, use_cache, return_audio_codes, **args) |
| tokens = list(self.stream_generate(input_ids, eos_token_id, max_new_tokens, temperature, top_p, rp, use_cache, return_audio_codes, **args)) |
| if tokens: |
| for text_out, _ in reversed(tokens): |
| if text_out is not None: |
| return text_out |
| return tokens[-1] |
| return input_ids |
|
|
| def stream_generate(self, input_ids, eos_token_id, max_new_tokens, temperature, top_p, rp, use_cache, return_audio_codes=False, **args): |
| start_pos, past_kvs, text_finished, first_finished = input_ids.shape[1], None, False, True |
| audio_codes = [[] for _ in range(8)] |
| audio_stop_pos = [None] * 8 |
| audio_buffer = torch.full((1, 8, start_pos), self.audio_pad_token, dtype=torch.long, device=input_ids.device) |
| spk_emb = args.get('spk_emb', None) |
| ref_codes = args.get('ref_codes', None) |
| ref_len = ref_codes.shape[2] if ref_codes is not None else 0 |
| spk_reserve = 1 if spk_emb is not None else 0 |
| fill_end = start_pos |
| fill_start = max(spk_reserve, start_pos - ref_len) |
| if ref_codes is not None and fill_start < fill_end: |
| audio_buffer[:, :, fill_start:fill_end] = ref_codes[:, :, -(fill_end - fill_start):] |
| if spk_emb is not None and fill_start > 0: |
| audio_buffer[:, :, fill_start - 1] = self.audio_spk_token |
| think_end_step, generated_tokens = None, ([] if args.get('open_thinking', False) else None) |
| while input_ids.shape[1] < start_pos + max_new_tokens: |
| if past_kvs is None or not use_cache: |
| out = self.forward(torch.cat((audio_buffer, input_ids.unsqueeze(1)), dim=1), past_key_values=past_kvs, use_cache=use_cache, **args) |
| else: |
| out = self.forward(torch.cat((audio_buffer[:, :, -1:], input_ids[:, -1:].unsqueeze(1)), dim=1), past_key_values=past_kvs, use_cache=use_cache, **args) |
| past_kvs = out.past_key_values |
|
|
| logits = out.logits[0, -1, :].clone().float() / (temperature + 1e-9) |
| logits = torch.nan_to_num(logits, nan=-100.0, posinf=-100.0, neginf=-100.0) |
| if rp != 1.0: |
| seen = list(set(input_ids[0].tolist())) |
| score = logits[seen] |
| logits[seen] = torch.where(score > 0, score / rp, score * rp) |
| if top_p and top_p < 1.0: |
| sorted_l, sorted_i = torch.sort(logits, descending=True) |
| mask = torch.cumsum(F.softmax(sorted_l, dim=-1), dim=-1) > top_p |
| mask[1:], mask[0] = mask[:-1].clone(), False |
| logits[sorted_i[mask]] = -float('Inf') |
| probs = F.softmax(logits, dim=-1) |
| probs = torch.nan_to_num(probs) |
| if probs.sum() <= 0: |
| probs = torch.ones_like(probs) / probs.shape[-1] |
| text_token = torch.multinomial(probs, 1).item() |
|
|
| if text_finished: |
| text_token = args.get('enter_token_id', 201) if first_finished else args.get('pad_token_id', 0) |
| first_finished = False |
|
|
| step = input_ids.shape[1] - start_pos |
| audio_step = step - 1 |
| if generated_tokens is not None: |
| generated_tokens.append(text_token) |
| if not think_end_step and generated_tokens[-len(self.config.think_end_ids):] == list(self.config.think_end_ids): |
| think_end_step = step + 2 |
| audio_step = (step - think_end_step) if think_end_step else -1 |
| for i, al in enumerate(out.audio_logits): |
| if audio_step < i: |
| audio_codes[i].append(self.audio_pad_token) |
| else: |
| logits_i = al[0, -1, :].clone().float() / 0.2 |
| logits_i = torch.nan_to_num(logits_i, nan=-100.0, posinf=-100.0, neginf=-100.0) |
| for prev_code in audio_codes[i][-3:]: |
| score = logits_i[prev_code] |
| logits_i[prev_code] = torch.where(score > 0, score / 1.05, score * 1.05) |
| top_val, top_idx = logits_i.topk(50) |
| probs = F.softmax(top_val, dim=-1) |
| probs = torch.nan_to_num(probs) |
| if probs.sum() <= 0: |
| probs = torch.ones_like(probs) / probs.shape[-1] |
| code = top_idx[torch.multinomial(probs, 1)].item() |
| audio_codes[i].append(code) |
| if audio_stop_pos[i] is None and code >= 2048: |
| audio_stop_pos[i] = len(audio_codes[i]) - 1 |
|
|
| if text_finished and all(audio_stop_pos[i] is not None for i in range(8)): |
| break |
|
|
| input_ids = torch.cat((input_ids, torch.tensor([[text_token]], device=input_ids.device)), dim=1) |
| audio_buffer = torch.cat((audio_buffer, torch.full((1, 8, 1), self.audio_pad_token, dtype=torch.long, device=input_ids.device)), dim=2) |
| for i in range(min(audio_step + 1, 8)): |
| audio_buffer[0, i, -1] = audio_codes[i][-1] |
|
|
| audio_frame = None |
| if return_audio_codes and audio_step >= 7: |
| frame = [audio_codes[i][step - 7 + i] for i in range(8)] |
| active_layers = sum(1 for i in range(8) if audio_stop_pos[i] is None or step - 7 + i < audio_stop_pos[i]) |
| if active_layers >= 8: |
| audio_frame = frame |
| if not text_finished: |
| yield input_ids[:, start_pos:], audio_frame |
| if text_token == eos_token_id: |
| text_finished = True |
| else: |
| yield None, audio_frame |
|
|