omni / src /models /vam /model.py
chenbhao's picture
Add NaN guard to audio code generation path
9422bfb
Raw
History Blame Contribute Delete
23.3 kB
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