xingweis's picture
Upload folder using huggingface_hub
edbe9c6 verified
Raw
History Blame Contribute Delete
4.96 kB
"""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 + "<think>\n\n</think>\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