"""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}