"""Dasheng + Qwen3 LLM model (inference only). Simplified version: no LoRA, no peft, no prompt_manager. Only 28 basic special tokens for audio generation. """ import os from typing import Any, Dict, List, Literal, Tuple import torch import torch.nn as nn from einops import rearrange from torch import Tensor from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer class DummyAudioEncoder(nn.Module): """Placeholder encoder for inference (audio encoder not used during generation).""" def __init__(self, embed_dim=768): super().__init__() self.embed_dim = embed_dim def forward(self, *args, **kwargs): raise RuntimeError("Audio encoder not available in inference-only mode") class AudioProjectorSubsample(nn.Module): def __init__(self, in_dim: int, out_dim: int, downsample_rate=5, hidden_dim: int | None = None): super().__init__() self.k = downsample_rate if hidden_dim is None: hidden_dim = out_dim self.net = nn.Sequential( nn.Linear(in_dim * self.k, hidden_dim), nn.GELU(), nn.Linear(hidden_dim, out_dim), ) def forward(self, x, mask=None): """ Downsample audio embeddings. :param x: [B, T, D] :param mask: [B, T] :return: [B, T', D'] """ batch_size, seq_len, dim = x.shape num_frames_to_discard = seq_len % self.k if num_frames_to_discard > 0: x = x[:, :-num_frames_to_discard, :] if mask is not None: mask = mask[:, :-num_frames_to_discard] if mask is None: mask = torch.ones(x.shape[:-1], dtype=torch.long, device=x.device) x = rearrange(x, "b (s k) d -> b s (k d)", k=self.k) x = self.net(x) mask = rearrange(mask, "b (s k) -> b s k", k=self.k) mask = mask.any(dim=-1).long() return x, mask class DashengQwen3Model(nn.Module): def __init__( self, audio_encoder: str = "DashengTokenizer", audio_encoder_args: Dict[str, Any] = dict(pretrained_from=None), text_model: str = "Qwen/Qwen3-1.7B", text_model_args: Dict[str, Any] = {}, subsample_factor: int = 5, use_encoderattention_mask: bool = True, system_prompt: str = "<|im_start|>user\n", prompt_right: str = "{task_prompt}<|im_end|>\n<|im_start|>assistant\n", disable_think: bool = True, **kwargs, ): super().__init__() self.subsample_factor = subsample_factor self.use_encoderattention_mask = use_encoderattention_mask self.audio_encoder_type = audio_encoder # Audio encoder (dummy for inference) self.audio_encoder = DummyAudioEncoder(embed_dim=768) # LLM decoder (load architecture only, weights from fine-tuned checkpoint) self.tokenizer = AutoTokenizer.from_pretrained(text_model) config = AutoConfig.from_pretrained(text_model) config.attn_implementation = "sdpa" self.decoder = AutoModelForCausalLM.from_config(config) # Add special tokens (must match training order for correct token IDs) special_tokens = [ "<|audio_bos|>", "<|audio_eos|>", "<|en|>", "<|kr|>", "<|de|>", "<|es|>", "<|fr|>", "<|hi|>", "<|uk|>", "<|th|>", "<|vi|>", "<|nl|>", "<|pt|>", "<|id|>", "<|ru|>", "<|it|>", "<|ar|>", "<|jp|>", "<|unknown|>", "<|AUDIO|>", "<|caption|>", "<|speech|>", "<|sfx|>", "<|music|>", "<|env|>", "<|asr|>", "<|speech_start|>", "<|speech_end|>", ] # old_vocab_size = len(self.tokenizer) self.tokenizer.add_special_tokens({"additional_special_tokens": special_tokens}) # new_vocab_size = len(self.tokenizer) # print(f"Added {new_vocab_size - old_vocab_size} special tokens, vocab: {old_vocab_size} -> {new_vocab_size}") self.decoder.resize_token_embeddings(len(self.tokenizer)) self.token_embed = self.decoder.model.embed_tokens # Audio projector self.audio_projector = AudioProjectorSubsample( self.audio_encoder.embed_dim, self.decoder.config.hidden_size, self.subsample_factor ) # Prompt templates self.system_prompt = system_prompt if disable_think: prompt_right = prompt_right + "\n\n\n\n" self.prompt_right = prompt_right self.audio_bos, self.audio_eos = "<|audio_bos|>", "<|audio_eos|>" def get_tokens(self, text, padding=False) -> Tuple[Tensor, Tensor]: """Tokenize text and return input_ids and attention_mask.""" tokens = self.tokenizer(text, add_special_tokens=False, padding=padding, return_tensors="pt") return tokens.input_ids.to(self.device), tokens.attention_mask.to(self.device) @property def device(self): return list(self.parameters())[0].device