import torch from transformers import PreTrainedModel, PretrainedConfig, GPT2TokenizerFast, \ AutoConfig, AutoModelForCausalLM from transformers.modeling_outputs import CausalLMOutput from model import FriendsTransformer class PerkLMConfig(PretrainedConfig): model_type = "perklm" def __init__(self, d_model = 512, n_heads = 8, n_layers = 6, d_ff = 2048, maxt = 512, dropout = 0.1, tokenizer_path = None, **kwargs): super().__init__(**kwargs) self.d_model = d_model self.n_heads = n_heads self.n_layers = n_layers self.d_ff = d_ff self.maxt = maxt self.dropout = dropout self.tokenizer_path = tokenizer_path class PerkLM(PreTrainedModel): config_class = PerkLMConfig _tied_weights_keys = ["transformer.lm_head.weight"] def __init__(self, config): super().__init__(config) tokenizer = GPT2TokenizerFast.from_pretrained(config.tokenizer_path) self.transformer = FriendsTransformer(d_model = config.d_model, n_heads = config.n_heads, n_layers = config.n_layers, d_ff = config.d_ff, dropout = config.dropout, maxt = config.maxt, tokenizer = tokenizer) self.post_init() def forward(self, input_ids, attention_mask = None, responder = None, labels = None, **kwargs): batch = {'input_ids': input_ids, 'attention_mask': attention_mask if attention_mask is not None \ else torch.ones_like(input_ids), 'responder': responder} logits = self.transformer(batch) loss = None if labels is not None: loss = torch.nn.functional.cross_entropy(logits[:, :-1].reshape(-1, logits.size(-1)), labels[:, 1:].reshape(-1), ignore_index = -100) return CausalLMOutput(loss = loss, logits = logits) def _save_pretrained_hook(self, *args, **kwargs): pass def state_dict(self, *args, **kwargs): sd = super().state_dict(*args, **kwargs) # strip complex64 RoPE buffers — recomputed on __init__ return {k: v for k, v in sd.items() if v.dtype != torch.complex64} def tie_weights(self): self.transformer.lm_head.weight = self.transformer.embedder.embedding.weight AutoConfig.register("perklm", PerkLMConfig) AutoModelForCausalLM.register(PerkLMConfig, PerkLM)