File size: 5,106 Bytes
769a891
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
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