Download model.py from VDC-team/VDrontV3-Mini: direct link, hf CLI and curl.
- Browser
- Download file 5.11 kB
-
https://huggingface.co/VDC-team/VDrontV3-Mini/resolve/main/model.py
- Command line
-
hf download hf://VDC-team/VDrontV3-Mini/model.py
-
curl -L -o model.py https://huggingface.co/VDC-team/VDrontV3-Mini/resolve/main/model.py
5.11 kB
| 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 | |
| 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 |