TechcodeX / modeling_hybrid.py
zarko1321's picture
Upload folder using huggingface_hub
49263e6 verified
Raw
History Blame Contribute Delete
2.57 kB
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}