| 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:
|
|
|
| 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
|
| )
|
|
|
| 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 |