import torch import torch.nn as nn from transformers import PreTrainedModel from .configuration_sai import SAIConfig from .generate import generate as _source_generate from .DecoderBlock import DecoderBlock as _SourceDecoderBlock from .GroupedQueryAttention import GroupedQueryAttention as _SourceGroupedQueryAttention from .RotaryPositionalEmbedding import RotaryPositionalEmbedding as _SourceRoPE from .SwiGLU import SwiGLU as _SourceSwiGLU from .TransformerModel import TransformerModel class SAIForCausalLM(PreTrainedModel): """Adapter mỏng giữa source SAI gốc và API tải model của Transformers.""" config_class = SAIConfig base_model_prefix = "sai" main_input_name = "input_ids" _tied_weights_keys = {"lm_head.weight": "embed.weight"} def __init__(self, config): PreTrainedModel.__init__(self, config) core = TransformerModel( vocab_size=config.vocab_size, d_model=config.d_model, num_heads=config.num_heads, num_kv_heads=config.num_kv_heads, num_layers=config.num_layers, ff_dim=config.ff_dim, max_seq_len=config.max_seq_len, dropout=config.dropout, pad_token_id=config.pad_token_id, ) # Giữ nguyên tên state_dict của checkpoint gốc, không thêm prefix "model.". self.d_model = core.d_model self.num_heads = core.num_heads self.num_kv_heads = core.num_kv_heads self.num_layers = core.num_layers self.pad_token_id = core.pad_token_id self.max_seq_len = core.max_seq_len self.embed = core.embed self.rope = core.rope self.blocks = core.blocks self.norm = core.norm self.lm_head = core.lm_head self.register_buffer("causal_mask", core.causal_mask, persistent=False) # Transformers 5.x tạo model trên device meta khi from_pretrained(). Hai # buffer non-persistent của source gốc cần được materialize ở lần chạy đầu. self._runtime_buffers_need_materialization = ( self.rope.cos_cached.is_meta or self.causal_mask.is_meta ) self.post_init() def _init_weights(self, module): std = self.d_model ** -0.5 if isinstance(module, (nn.Linear, nn.Embedding)): nn.init.normal_(module.weight, mean=0.0, std=std) if isinstance(module, nn.Linear) and module.bias is not None: nn.init.zeros_(module.bias) def _ensure_runtime_buffers(self, device): if not self._runtime_buffers_need_materialization: return head_dim = self.d_model // self.num_heads inv_freq = 1.0 / ( 10_000 ** ( torch.arange(0, head_dim, 2, dtype=torch.float32, device=device) / head_dim ) ) positions = torch.arange(self.max_seq_len, dtype=torch.float32, device=device) frequencies = torch.outer(positions, inv_freq) embedding = torch.cat([frequencies, frequencies], dim=-1) self.rope.register_buffer("inv_freq", inv_freq, persistent=False) self.rope.register_buffer("cos_cached", embedding.cos(), persistent=False) self.rope.register_buffer("sin_cached", embedding.sin(), persistent=False) causal = torch.triu( torch.full( (self.max_seq_len, self.max_seq_len), float("-inf"), device=device, ), diagonal=1, ) self.register_buffer("causal_mask", causal, persistent=False) self._runtime_buffers_need_materialization = False def get_input_embeddings(self): return self.embed def set_input_embeddings(self, value): self.embed = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, value): self.lm_head = value def _build_attn_mask(self, sequence_length, pad_mask, device): self._ensure_runtime_buffers(device) return TransformerModel._build_attn_mask( self, sequence_length, pad_mask, device ) def forward_features(self, input_ids, attention_mask=None, has_padding=True): self._ensure_runtime_buffers(input_ids.device) return TransformerModel.forward_features( self, input_ids, attention_mask, has_padding ) def forward(self, input_ids, attention_mask=None, has_padding=True, **kwargs): return TransformerModel.forward( self, input_ids, attention_mask, has_padding ) def init_cache(self, batch_size, max_gen_len, device): return TransformerModel.init_cache(self, batch_size, max_gen_len, device) def prefill(self, input_ids, kv_cache=None): self._ensure_runtime_buffers(input_ids.device) return TransformerModel.prefill(self, input_ids, kv_cache) def decode_step(self, token_ids, kv_cache, cache_len): self._ensure_runtime_buffers(token_ids.device) return TransformerModel.decode_step(self, token_ids, kv_cache, cache_len) def generate_response(self, user_input, tokenizer, **kwargs): sentencepiece = getattr(tokenizer, "sp_model", tokenizer) return TransformerModel.generate_response( self, user_input, sentencepiece, **kwargs ) def generate(self, user_input, tokenizer, **kwargs): return self.generate_response(user_input, tokenizer, **kwargs)