| import torch |
| import torch.nn as nn |
| from torch.nn import functional as F |
| from transformers import PretrainedConfig, PreTrainedModel |
|
|
| class TechcodeXConfig(PretrainedConfig): |
| model_type = "techcodex_hybrid_transformer" |
| def __init__(self, vocab_size=50257, n_embd=256, n_head=4, block_size=128, **kwargs): |
| super().__init__(**kwargs) |
| self.vocab_size = vocab_size |
| self.n_embd = n_embd |
| self.n_head = n_head |
| self.block_size = block_size |
|
|
| class ProprietaryRecurrentLayer(nn.Module): |
| def __init__(self, n_embd): |
| super().__init__() |
| self.hidden_dim = n_embd |
| self.gate_mix = nn.Linear(n_embd * 2, n_embd) |
| self.gate_state = nn.Linear(n_embd * 2, n_embd) |
| self.ln = nn.LayerNorm(n_embd) |
|
|
| def forward(self, x): |
| B, T, C = x.shape |
| hidden = torch.zeros(B, self.hidden_dim, device=x.device) |
| outputs = [] |
| for t in range(T): |
| current_token = x[:, t, :] |
| combined = torch.cat([current_token, hidden], dim=-1) |
| mix = torch.sigmoid(self.gate_mix(combined)) |
| new_state = torch.tanh(self.gate_state(combined)) |
| hidden = (mix * hidden) + ((1.0 - mix) * new_state) |
| outputs.append(hidden.unsqueeze(1)) |
| return self.ln(torch.cat(outputs, dim=1)) |
|
|
| class TechcodeXModel(PreTrainedModel): |
| config_class = TechcodeXConfig |
| def __init__(self, config): |
| super().__init__(config) |
| self.config = config |
| self.token_embedding = nn.Embedding(config.vocab_size, config.n_embd) |
| self.position_embedding = nn.Embedding(config.block_size, config.n_embd) |
| self.attn_layer = nn.TransformerEncoderLayer( |
| d_model=config.n_embd, nhead=config.n_head, dim_feedforward=config.n_embd*4, batch_first=True |
| ) |
| self.recurrent_layer = ProprietaryRecurrentLayer(config.n_embd) |
| self.bridge = nn.Linear(config.n_embd * 2, config.n_embd) |
| self.ln_final = nn.LayerNorm(config.n_embd) |
| self.lm_head = nn.Linear(config.n_embd, config.vocab_size) |
| self.post_init() |
|
|
| def forward(self, input_ids, labels=None, **kwargs): |
| B, T = input_ids.shape |
| positions = torch.arange(0, T, device=input_ids.device).unsqueeze(0) |
| x = self.token_embedding(input_ids) + self.position_embedding(positions) |
| path_a = self.attn_layer(x) |
| path_b = self.recurrent_layer(x) |
| x = self.bridge(torch.cat([path_a, path_b], dim=-1)) |
| logits = self.lm_head(self.ln_final(x)) |
| return {"logits": logits} |
|
|