import torch import torch.nn as nn import torch.nn.functional as F from transformers import GPT2Config from transformers.models.gpt2.modeling_gpt2 import GPT2Block from typing import List import copy class ExpertBlock(nn.Module): def __init__(self, layers: List[nn.Module], num_versions: int, config=None): super().__init__() self.num_versions = num_versions self.config = config self.versions = nn.ModuleList([ nn.ModuleList([copy.deepcopy(layer) for layer in layers]) for _ in range(num_versions) ]) self.active_version = 0 def set_version(self, idx): self.active_version = idx def forward(self, x, **kwargs): for layer in self.versions[self.active_version]: out = layer(x, **kwargs) if isinstance(out, tuple): x = out[0] else: x = out return x class VDrontModel(nn.Module): def __init__(self, config, expert_start, expert_end, output_index, num_experts, num_output_versions): super().__init__() self.config = config self.num_experts = num_experts self.num_output_versions = num_output_versions self.expert_start = expert_start self.expert_end = expert_end self.output_index = output_index self.embed_tokens = nn.Embedding(config.vocab_size, config.n_embd) self.embed_positions = nn.Embedding(config.n_positions, config.n_embd) all_layers = [GPT2Block(config, layer_idx=i) for i in range(config.n_layer)] expert_layers = all_layers[expert_start:expert_end+1] base_layers = all_layers[expert_end+1:output_index] output_layer = all_layers[output_index] self.expert_block = ExpertBlock(expert_layers, num_experts, config=config) self.base_blocks = nn.ModuleList(base_layers) self.output_block = ExpertBlock([output_layer], num_output_versions, config=config) self.ln_f = nn.LayerNorm(config.n_embd, eps=config.layer_norm_epsilon) self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False) self.router = nn.Linear(config.n_embd, num_experts, bias=False) self.apply(self._init_weights) def _init_weights(self, module): if isinstance(module, (nn.Linear, nn.Embedding)): module.weight.data.normal_(mean=0.0, std=0.02) if isinstance(module, nn.Linear) and module.bias is not None: module.bias.data.zero_() def set_expert_version(self, idx): self.expert_block.set_version(idx) def set_output_version(self, idx): self.output_block.set_version(idx) def forward(self, input_ids, labels=None, return_router_logits=False): pos = torch.arange(0, input_ids.size(1), device=input_ids.device).unsqueeze(0) x = self.embed_tokens(input_ids) + self.embed_positions(pos) router_logits = self.router(x.mean(dim=1)) if return_router_logits else None x = self.expert_block(x) for block in self.base_blocks: out = block(x) x = out[0] if isinstance(out, tuple) else out x = self.output_block(x) x = self.ln_f(x) logits = self.lm_head(x) loss = None if labels is not None: # Защита: всё, что вне [0, vocab_size), заменяем на -100 (игнорируем) labels = torch.where( (labels >= 0) & (labels < self.config.vocab_size), labels, -100 ) loss = F.cross_entropy( logits.reshape(-1, logits.size(-1)), labels.reshape(-1), ignore_index=-100 # ВАЖНО: именно -100, а не -1 ) if return_router_logits: return logits, loss, router_logits return logits, loss @torch.no_grad() def generate(self, input_ids, max_new_tokens, temperature=1.0, top_k=None, dynamic_expert=True): self.eval() for _ in range(max_new_tokens): if dynamic_expert: pos = torch.arange(0, input_ids.size(1), device=input_ids.device).unsqueeze(0) x = self.embed_tokens(input_ids) + self.embed_positions(pos) router_logits = self.router(x.mean(dim=1)) expert_idx = router_logits.argmax(dim=-1).item() self.set_expert_version(expert_idx) idx_cond = input_ids[:, -self.config.n_positions:] logits, _ = self(idx_cond) logits = logits[:, -1, :] / temperature if top_k is not None: v, _ = torch.topk(logits, min(top_k, logits.size(-1))) logits[logits < v[:, [-1]]] = -float('Inf') probs = F.softmax(logits, dim=-1) idx_next = torch.multinomial(probs, num_samples=1) input_ids = torch.cat((input_ids, idx_next), dim=1) return input_ids