pc-mlp-tiny / modeling_pcmlp.py
zeechimp's picture
Create modeling_pcmlp.py
491a99b verified
Raw History Blame Contribute Delete
1.69 kB
"""PC-MLP model for HuggingFace Hub with auto_map."""
import math
import torch
import torch.nn as nn
from transformers import PreTrainedModel
from .configuration_pcmlp import PCMLPConfig
class PCMLP(PreTrainedModel):
config_class = PCMLPConfig
def __init__(self, config):
super().__init__(config)
d = config.vocab_size
h = config.hidden_dim
scale = 1.0 / math.sqrt(h)
self.tok = nn.Embedding(config.vocab_size, d)
self.pos = nn.Embedding(config.seq_len, d)
self.W1 = nn.Linear(d, h, bias=False)
self.W2 = nn.Linear(h, h, bias=False)
self.head = nn.Linear(h, config.num_classes)
# match JAX init scale
nn.init.normal_(self.tok.weight, std=0.05)
nn.init.normal_(self.pos.weight, std=0.05)
nn.init.normal_(self.W1.weight, std=scale)
nn.init.normal_(self.W2.weight, std=scale)
nn.init.normal_(self.head.weight, std=0.02)
def forward(self, input_ids, labels=None):
B, T = input_ids.shape
pos_ids = torch.arange(T, device=input_ids.device)
h = self.tok(input_ids) + self.pos(pos_ids)
h = h.reshape(B, -1)[:, :self.W1.in_features]
h0 = h
h1 = torch.nn.functional.gelu(self.W1(h0))
h2 = torch.nn.functional.gelu(self.W2(h1))
logits = self.head(h2)
loss = None
if labels is not None:
ce = torch.nn.functional.cross_entropy(logits, labels)
local = ((h1.mean(1) - h0.mean(1)) ** 2).mean() + \
((h2.mean(1) - h1.mean(1)) ** 2).mean()
loss = ce + self.config.local_loss_weight * local
return {"loss": loss, "logits": logits}