"""Audio-conditioned Thinker and native-token speech generation.""" from __future__ import annotations from typing import Iterator import torch from torch import nn from torch.nn.attention import SDPBackend, sdpa_kernel from transformers import GenerationMixin, PreTrainedModel, Qwen3_5ForConditionalGeneration from transformers.modeling_outputs import CausalLMOutputWithPast from transformers.models.qwen3_omni_moe.modeling_qwen3_omni_moe import ( Qwen3OmniMoeAudioEncoder, _get_feat_extract_output_lengths, ) from .configuration_edgeinstant import EdgeInstantConfig from .codec_edgeinstant import EdgeInstantCodecDecoder from .inference_edgeinstant import EdgeInstantDecoderLayer from .compact_native_clock import CompactNativeClockTalker class EdgeInstantThinker(Qwen3_5ForConditionalGeneration): """Qwen3.5 with CUDA graph replay for cached recurrent decoding.""" _no_split_modules = ["Qwen3_5DecoderLayer", "EdgeInstantDecoderLayer", "Qwen3_5VisionBlock"] def __init__(self, config): super().__init__(config) for index in range(len(self.model.language_model.layers)): if config.text_config.layer_types[index] == "linear_attention": self.model.language_model.layers[index] = EdgeInstantDecoderLayer(config.text_config, index) class EdgeInstantAudioProjector(nn.Module): def __init__(self, config: dict, hidden_size: int): super().__init__() self.hidden_size = hidden_size self.mode = config.get("align_mode", "replace") if self.mode == "residual" and int(config["output_size"]) != hidden_size: raise ValueError("Residual audio projection requires audio_split=1 (output_size must equal Thinker hidden_size)") self.out_norm = bool(config.get("audio_out_norm", False)) self.norm = nn.LayerNorm(int(config["norm_size"])) middle = config.get("intermediate_size") self.proj = nn.Linear(int(config["input_size"]), int(middle or config["output_size"])) self.proj2 = nn.Linear(int(middle), int(config["output_size"])) if middle else None self.gate = nn.Parameter(torch.zeros(())) self.out_scale = nn.Parameter(torch.ones(())) def forward(self, audio: torch.Tensor, pad_embedding: torch.Tensor) -> torch.Tensor: shape = audio.shape frames = audio.float().reshape(*shape[:-1], -1, self.norm.normalized_shape[0]) mapped = self.proj(self.norm(frames).reshape(shape)) if self.proj2 is not None: mapped = self.proj2(nn.functional.gelu(mapped)) if self.mode == "residual": return pad_embedding.to(mapped.dtype) + self.gate.clamp(0, 4) * mapped if self.out_norm: mapped = nn.functional.normalize(mapped, dim=-1, eps=1e-6) * self.hidden_size ** 0.5 return (mapped * self.out_scale).reshape(-1, self.hidden_size) class EdgeInstantForConditionalGeneration(PreTrainedModel, GenerationMixin): config_class = EdgeInstantConfig base_model_prefix = "edgeinstant" main_input_name = "input_ids" _supports_sdpa = True _keep_in_fp32_modules_strict = ["align", "audio_special"] _no_split_modules = ["Qwen3OmniMoeAudioEncoder", "EdgeInstantAudioProjector", "EdgeInstantCodecDecoder"] def __init__(self, config, *, thinker=None, audio_tower=None, align=None, talker=None, codec=None, audio_special_delta=None): super().__init__(config) supplied = [component for component in (thinker, audio_tower, align, talker, codec) if component is not None] self.thinker = thinker if thinker is not None else EdgeInstantThinker(config.thinker_config) self.audio_tower = audio_tower if audio_tower is not None else Qwen3OmniMoeAudioEncoder(config.audio_config) if config.audio_tap == "ln_post": self.audio_tower.proj1 = nn.Identity() self.audio_tower.act = nn.Identity() self.audio_tower.proj2 = nn.Identity() self.align = align if align is not None else EdgeInstantAudioProjector( config.projector_config, config.thinker_config.text_config.hidden_size, ) self.audio_special = nn.Embedding(len(config.audio_special_ids), config.thinker_config.text_config.hidden_size) if audio_special_delta is not None: self.audio_special.weight = nn.Parameter(audio_special_delta.detach().clone()) supplied.append(self.audio_special) self.talker = talker if talker is not None else CompactNativeClockTalker.from_config(config.talker_config) self.talker.prepare_for_serialization() self.codec = codec if codec is not None else ( EdgeInstantCodecDecoder(config.codec_config) if config.codec_config else None ) self._tied_weights_keys = { f"talker.{target}": f"talker.{source}" for target, source in self.talker.tied_weight_keys.items() } for component in supplied: for module in component.modules(): module._is_hf_initialized = True self.post_init() def get_input_embeddings(self): return self.thinker.get_input_embeddings() def set_input_embeddings(self, embeddings): self.thinker.set_input_embeddings(embeddings) def get_output_embeddings(self): return self.thinker.get_output_embeddings() def _base_thinker(self): return self.thinker.get_base_model() if hasattr(self.thinker, "get_base_model") else self.thinker def encode_audio(self, input_features, feature_attention_mask=None, tower=None, feature_lengths=None): """Return packed audio encoder features before the trainable projector.""" from transformers.models.qwen3_omni_moe.modeling_qwen3_omni_moe import ( chunk_and_pad_features, get_audio_cu_seqlens, get_valid_indices, ) from transformers.utils.generic import is_flash_attention_requested tower = self.audio_tower if tower is None else tower tower_weight = next(tower.parameters()) if feature_lengths is not None: lengths = torch.as_tensor(feature_lengths, dtype=torch.long).cpu() elif feature_attention_mask is not None: lengths = feature_attention_mask.sum(-1).to(device="cpu", dtype=torch.long) else: lengths = torch.full((input_features.shape[0],), input_features.shape[-1], dtype=torch.long) if bool((lengths <= 0).any()): raise ValueError("Audio must contain at least one valid mel frame") packed = torch.cat([row[:, :int(length)] for row, length in zip(input_features, lengths)], dim=-1) packed = packed.to(device=tower_weight.device, dtype=tower_weight.dtype) padded_feature, chunk_lengths = chunk_and_pad_features(packed, lengths, tower.n_window) valid_indices = get_valid_indices(chunk_lengths).to(tower_weight.device) cu_seqlens = get_audio_cu_seqlens(chunk_lengths, lengths, tower.n_window_infer, tower.n_window) # SDPA/eager split each attention window in Python; keep their boundaries on the CPU. if is_flash_attention_requested(tower.config): cu_seqlens = cu_seqlens.to(tower_weight.device) hidden = tower( input_features=packed, feature_lens=lengths, padded_feature=padded_feature, chunk_lengths=chunk_lengths, valid_indices=valid_indices, cu_seqlens=cu_seqlens, ).last_hidden_state return self.group_audio_features(hidden, lengths) def group_audio_features(self, hidden, frame_lengths): """Apply per-recording frame stacking or pooling to packed encoder output.""" if self.config.audio_stack == 1 and self.config.audio_pool == 1: return hidden counts = _get_feat_extract_output_lengths(frame_lengths).tolist() processed = [] for row in hidden.split(counts): if self.config.audio_stack > 1: padding = -len(row) % self.config.audio_stack if padding: row = torch.cat((row, row[-1:].expand(padding, -1))) row = row.reshape(len(row) // self.config.audio_stack, -1) elif self.config.audio_pool > 1: row = torch.stack([chunk.mean(0) for chunk in row.split(self.config.audio_pool)]) processed.append(row) return torch.cat(processed) def pad_embedding(self): """Audio placeholder embedding including its learned special-token delta.""" embed = self.get_input_embeddings() pad = embed.weight[self.config.audio_pad_id] if self.config.audio_pad_id in self.config.audio_special_ids: slot = self.config.audio_special_ids.index(self.config.audio_pad_id) pad = pad + self.audio_special.weight[slot].to(pad) return pad def get_audio_features(self, input_features, feature_attention_mask, feature_lengths=None): features = self.encode_audio(input_features, feature_attention_mask, feature_lengths=feature_lengths) device = self.align.proj.weight.device return self.align(features.to(device), self.pad_embedding().to(device)) def merge_audio_embeds(self, input_ids, input_features, feature_attention_mask, feature_lengths=None): vectors = self.get_audio_features(input_features, feature_attention_mask, feature_lengths) return self.merge_custom_audio(input_ids, vectors) def merge_custom_audio(self, input_ids, audio_vectors): """Insert projected or diagnostic audio vectors in the model's audio slots.""" embedding = self.get_input_embeddings() ids = input_ids.to(embedding.weight.device) embeds = embedding(ids) for slot, token_id in enumerate(self.config.audio_special_ids): embeds = embeds + (ids == token_id).unsqueeze(-1) * self.audio_special.weight[slot].to(embeds) vectors = audio_vectors.to(embeds) mask = ids == self.config.audio_pad_id if int(mask.sum()) != len(vectors): raise ValueError(f"Audio slots {int(mask.sum())} do not match encoded vectors {len(vectors)}") return embeds.masked_scatter(mask.unsqueeze(-1).expand_as(embeds), vectors) def forward(self, input_ids=None, attention_mask=None, input_features=None, feature_attention_mask=None, inputs_embeds=None, labels=None, audio_feature_lengths=None, answer_loss_reduction=None, eos_loss_weight=1.0, unlikelihood_mask=None, unlikelihood_weight=1.0, **kwargs): if input_features is not None: if input_ids is None or feature_attention_mask is None: raise ValueError("Audio forward requires input_ids and feature_attention_mask") cache = kwargs.get("past_key_values") if cache is None or cache.get_seq_length() == 0: self._base_thinker().model.rope_deltas = None inputs_embeds = self.merge_audio_embeds( input_ids, input_features, feature_attention_mask, audio_feature_lengths, ) if kwargs.get("position_ids") is None and ( kwargs.get("image_grid_thw") is not None or kwargs.get("video_grid_thw") is not None ): kwargs["position_ids"] = self._base_thinker().model.compute_3d_position_ids( input_ids, inputs_embeds, attention_mask=attention_mask, **{key: kwargs.get(key) for key in ("image_grid_thw", "video_grid_thw", "mm_token_type_ids")}, ) input_ids = None if answer_loss_reduction is not None: if labels is None or answer_loss_reduction not in {"token_sum", "example_sum"}: raise ValueError("Answer loss requires labels and token_sum or example_sum reduction") # LoRA lives on the base model's decoder linears; the vocabulary projection # only needs hidden states that predict a supervised answer token or EOS. thinker = self._base_thinker() outputs = thinker.model(input_ids=input_ids, inputs_embeds=inputs_embeds, attention_mask=attention_mask, **kwargs) targets = labels[:, 1:] mask = targets.ne(-100) hidden = outputs.last_hidden_state[:, :-1][mask] logits = thinker.lm_head(hidden) losses = nn.functional.cross_entropy(logits.float(), targets[mask], reduction="none") if unlikelihood_mask is not None: negative = unlikelihood_mask[:, 1:][mask].bool() if negative.any(): negative_logits = logits[negative].float() negative_ids = targets[mask][negative, None] selected = negative_logits.gather(1, negative_ids).squeeze(1) alternatives = negative_logits.scatter(1, negative_ids, -torch.inf).logsumexp(dim=1) losses = losses.index_put((negative,), nn.functional.softplus(selected - alternatives) * unlikelihood_weight) eos_ids = self.config.eos_token_id if isinstance(eos_ids, int): eos_ids = [eos_ids] is_eos = torch.isin(targets[mask], targets.new_tensor(eos_ids)) weights = torch.where(is_eos, eos_loss_weight, 1.0) losses = losses * weights if answer_loss_reduction == "example_sum": example_ids = mask.nonzero(as_tuple=True)[0] counts = weights.new_zeros(targets.shape[0]).scatter_add(0, example_ids, weights) losses = losses / counts[example_ids] return CausalLMOutputWithPast(loss=losses.sum()) return self.thinker(input_ids=input_ids, inputs_embeds=inputs_embeds, attention_mask=attention_mask, labels=labels, **kwargs) @torch.no_grad() @sdpa_kernel([SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION, SDPBackend.MATH]) def generate(self, input_ids=None, attention_mask=None, input_features=None, feature_attention_mask=None, **kwargs): """Generate text token IDs, including the supplied prompt as in HF causal LMs.""" self._base_thinker().model.rope_deltas = None if input_features is not None: kwargs["inputs_embeds"] = self.merge_audio_embeds(input_ids, input_features, feature_attention_mask) self.thinker.generation_config = self.generation_config kwargs.setdefault("eos_token_id", self.config.eos_token_id) return self.thinker.generate(input_ids=input_ids, attention_mask=attention_mask, **kwargs) generate_text = generate @torch.inference_mode() def synthesize(self, token_ids, *, language="chinese", max_frames=1536, do_sample=False, decode=True, **sampling): """Speak one sequence of native Thinker token IDs.""" state = self.talker.new_state(language=language, max_frames=max_frames, do_sample=do_sample, **sampling) for token in torch.as_tensor(token_ids).reshape(-1).tolist(): state.push_token(token) state.end_input() while not state.ended: state.advance(max_new_frames=4) codes = torch.stack(state.frames) if state.frames else torch.empty((0, 16), dtype=torch.long) waveform = None if decode: if self.codec is None: raise ValueError("The model package does not contain a codec decoder") waveform = self.codec.decode(codes.unsqueeze(0))[0] if len(codes) else torch.empty(0) return {"audio_codes": codes, "audio": waveform, "sampling_rate": self.config.output_sampling_rate, "stop_reason": state.status.value} @torch.inference_mode() def generate_speech(self, input_ids=None, *, language="chinese", max_frames=1536, audio_do_sample=False, **kwargs): """Generate one reply and its waveform; text generation accepts standard HF options.""" if input_ids is None or input_ids.shape[0] != 1: raise ValueError("Speech generation accepts one conversation at a time") sequences = self.generate(input_ids=input_ids, **kwargs) if not isinstance(sequences, torch.Tensor): sequences = sequences.sequences reply = sequences[0, input_ids.shape[1]:] special = set(self.config.thinker_special_ids) public, thinking = [], False for token in reply.tolist(): if token == self.config.think_start_id: thinking = True elif token == self.config.think_end_id: thinking = False elif not thinking and token not in special: public.append(token) result = self.synthesize(public, language=language, max_frames=max_frames, do_sample=audio_do_sample) return {"sequences": sequences, "text_token_ids": torch.tensor(public), **result} @torch.inference_mode() def stream_generate(self, input_ids, attention_mask=None, input_features=None, feature_attention_mask=None, *, language="chinese", max_new_tokens=512, max_frames=1536, first_packet_frames=1, chunk_frames=4, left_context_frames=50, audio_do_sample=False, cancelled=None) -> Iterator[dict]: """Yield native text tokens and continuous 24 kHz audio chunks for one reply.""" if input_ids.shape[0] != 1: raise ValueError("Streaming accepts one conversation at a time") state = self.talker.new_state(language=language, max_frames=max_frames, do_sample=audio_do_sample) self._base_thinker().model.rope_deltas = None inputs_embeds = None if input_features is not None: inputs_embeds = self.merge_audio_embeds(input_ids, input_features, feature_attention_mask) if attention_mask is None: attention_mask = torch.ones_like(input_ids) sequences = input_ids sent = 0 special = set(self.config.thinker_special_ids) eos = self.config.eos_token_id eos = {eos} if isinstance(eos, int) else set(eos) thinking = False def packet(): nonlocal sent if len(state.frames) == sent: return None start = max(0, sent - left_context_frames) codes = torch.stack(state.frames[start:]).unsqueeze(0) pcm = self.codec(codes)[0] pcm = pcm[(sent - start) * self.codec.samples_per_frame:] event = {"type": "audio", "audio": pcm.detach().cpu(), "sampling_rate": self.codec.sample_rate, "sample_start": sent * self.codec.samples_per_frame, "audio_codes": torch.stack(state.frames[sent:])} sent = len(state.frames) return event if self.codec is None: raise ValueError("The model package does not contain a codec decoder") try: inputs = self.thinker.prepare_inputs_for_generation( sequences, inputs_embeds=inputs_embeds, attention_mask=attention_mask, is_first_iteration=True, use_cache=True, logits_to_keep=1, ) output = self.thinker(**inputs) cache = output.past_key_values reached_eos = False for _ in range(max_new_tokens): if cancelled is not None and cancelled.is_set(): return token = output.logits[:, -1].argmax(-1) token_id = int(token.item()) if token_id in eos: reached_eos = True break if token_id == self.config.think_start_id: thinking = True elif token_id == self.config.think_end_id: thinking = False elif not thinking and token_id not in special: yield {"type": "text", "token_id": token_id} state.push_token(token_id) state.advance(max_new_frames=first_packet_frames if sent == 0 else chunk_frames) event = packet() if event is not None: yield event if state.ended: break sequences = torch.cat((sequences, token[:, None]), dim=-1) attention_mask = torch.cat((attention_mask, torch.ones_like(token[:, None])), dim=-1) inputs = self.thinker.prepare_inputs_for_generation( sequences, past_key_values=cache, attention_mask=attention_mask, is_first_iteration=False, next_sequence_length=1, use_cache=True, logits_to_keep=1, ) output = self.thinker(**inputs) cache = output.past_key_values if not state.ended: if not reached_eos: raise RuntimeError(f"Thinker did not reach EOS within {max_new_tokens} tokens") state.end_input() while not state.ended: if cancelled is not None and cancelled.is_set(): return state.advance(max_new_frames=chunk_frames) event = packet() if event is not None: yield event yield {"type": "done", "stop_reason": state.status.value, "samples": sent * self.codec.samples_per_frame} finally: state.cancel()