VDrontV3-Mini / model.py
MishaGGG's picture
Upload 12 files
769a891 verified
Raw
History Blame Contribute Delete
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
@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