| """ |
| PIT (Point-In-Time) GPT model — self-contained for trust_remote_code=True loading. |
| |
| Architecture: decoder-only Transformer with RoPE, RMSNorm on Q/K, squared-ReLU |
| MLP, and weight-tied input/output embeddings. |
| |
| Generation is KV-cached: past keys and values are kept in a `DynamicCache`, so |
| each step runs attention with a single query position against the cache instead |
| of re-running the whole prefix. |
| """ |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from transformers import GenerationMixin, PreTrainedModel |
| from transformers.cache_utils import Cache, DynamicCache |
| from transformers.modeling_outputs import CausalLMOutputWithPast |
|
|
| from .configuration_pit import PITConfig |
|
|
|
|
| |
| |
| |
|
|
| class Rotary(nn.Module): |
| def __init__(self, dim: int, base: int = 10000, scaling_factor: float = 1.0): |
| super().__init__() |
| self.dim = dim |
| self.base = base * scaling_factor |
| self.inv_freq: torch.Tensor | None = None |
|
|
| def forward(self, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: |
| """Absolute positions [B, T] -> (cos, sin), each of shape [B, T, 1, dim // 2]. |
| |
| Positions are passed in rather than derived from the sequence length so |
| that a cached decode step rotates the new token by its *absolute* |
| position, not by 0. |
| """ |
| if self.inv_freq is None or self.inv_freq.device != position_ids.device: |
| |
| |
| self.inv_freq = 1.0 / (self.base ** ( |
| torch.arange(0, self.dim, 2, device=position_ids.device, dtype=torch.float32) / self.dim |
| )) |
| freqs = position_ids.float()[:, :, None] * self.inv_freq |
| return freqs.cos().bfloat16()[:, :, None, :], freqs.sin().bfloat16()[:, :, None, :] |
|
|
|
|
| def _apply_rotary_emb(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: |
| d = x.shape[3] // 2 |
| x1, x2 = x[..., :d], x[..., d:] |
| return torch.cat([x1 * cos + x2 * sin, x1 * (-sin) + x2 * cos], dim=3).type_as(x) |
|
|
|
|
| class CausalSelfAttention(nn.Module): |
| def __init__(self, config: PITConfig, layer_idx: int = 0): |
| super().__init__() |
| self.layer_idx = layer_idx |
| self.n_head = config.n_head |
| self.n_embd = config.n_embd |
| self.head_dim = config.n_embd // config.n_head |
| self.c_q = nn.Linear(config.n_embd, config.n_embd, bias=False) |
| self.c_k = nn.Linear(config.n_embd, config.n_embd, bias=False) |
| self.c_v = nn.Linear(config.n_embd, config.n_embd, bias=False) |
| self.c_proj = nn.Linear(config.n_embd, config.n_embd, bias=False) |
| self.c_proj.weight.data.zero_() |
|
|
| def forward( |
| self, |
| x: torch.Tensor, |
| cos: torch.Tensor, |
| sin: torch.Tensor, |
| attn_mask: torch.Tensor | None = None, |
| is_causal: bool = False, |
| past_key_values: Cache | None = None, |
| ) -> torch.Tensor: |
| B, T, C = x.size() |
| q = self.c_q(x).view(B, T, self.n_head, self.head_dim) |
| k = self.c_k(x).view(B, T, self.n_head, self.head_dim) |
| v = self.c_v(x).view(B, T, self.n_head, self.head_dim) |
| q = _apply_rotary_emb(F.rms_norm(q, (q.size(-1),)), cos, sin) |
| k = _apply_rotary_emb(F.rms_norm(k, (k.size(-1),)), cos, sin) |
|
|
| |
| q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2) |
| if past_key_values is not None: |
| |
| k, v = past_key_values.update(k, v, self.layer_idx) |
|
|
| y = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, is_causal=is_causal) |
| return self.c_proj(y.transpose(1, 2).contiguous().view_as(x)) |
|
|
|
|
| class MLP(nn.Module): |
| def __init__(self, config: PITConfig): |
| super().__init__() |
| self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd, bias=False) |
| self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd, bias=False) |
| self.c_proj.weight.data.zero_() |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| return self.c_proj(F.relu(self.c_fc(x)).square()) |
|
|
|
|
| class Block(nn.Module): |
| def __init__(self, config: PITConfig, layer_idx: int = 0): |
| super().__init__() |
| self.attn = CausalSelfAttention(config, layer_idx) |
| self.mlp = MLP(config) |
|
|
| def forward( |
| self, |
| x: torch.Tensor, |
| cos: torch.Tensor, |
| sin: torch.Tensor, |
| attn_mask: torch.Tensor | None = None, |
| is_causal: bool = False, |
| past_key_values: Cache | None = None, |
| ) -> torch.Tensor: |
| x = x + self.attn( |
| F.rms_norm(x, (x.size(-1),)), cos, sin, attn_mask, is_causal, past_key_values |
| ) |
| x = x + self.mlp(F.rms_norm(x, (x.size(-1),))) |
| return x |
|
|
|
|
| |
| |
| |
|
|
| class PITForCausalLM(PreTrainedModel, GenerationMixin): |
| """ |
| Point-In-Time GPT wrapped as a HuggingFace CausalLM. |
| |
| Supports AutoModelForCausalLM, generate(), and pipeline("text-generation"). |
| |
| Loading |
| ------- |
| >>> from transformers import AutoTokenizer, AutoModelForCausalLM |
| >>> tokenizer = AutoTokenizer.from_pretrained("Diamegs/PIT-4B-FT-202012") |
| >>> model = AutoModelForCausalLM.from_pretrained( |
| ... "Diamegs/PIT-4B-FT-202012", |
| ... trust_remote_code=True, |
| ... torch_dtype=torch.bfloat16, |
| ... device_map="auto", |
| ... ) |
| """ |
|
|
| config_class = PITConfig |
| _no_split_modules = ["Block"] |
| _supports_cache_class = True |
| |
| |
| _can_compile_fullgraph = False |
| |
| |
| |
| _tied_weights_keys = {"lm_head.weight": "transformer.wte.weight"} |
|
|
| def __init__(self, config: PITConfig): |
| super().__init__(config) |
| self.transformer = nn.ModuleDict({ |
| "wte": nn.Embedding(config.vocab_size, config.n_embd), |
| "h": nn.ModuleList([Block(config, i) for i in range(config.n_layer)]), |
| }) |
| self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False) |
| |
| |
| self.rotary = Rotary(config.n_embd // config.n_head) |
| |
| self.transformer["wte"].weight = self.lm_head.weight |
| self.post_init() |
|
|
| |
|
|
| def get_input_embeddings(self) -> nn.Embedding: |
| return self.transformer["wte"] |
|
|
| def set_input_embeddings(self, value: nn.Embedding) -> None: |
| self.transformer["wte"] = value |
|
|
| def get_output_embeddings(self) -> nn.Linear: |
| return self.lm_head |
|
|
| def set_output_embeddings(self, value: nn.Linear) -> None: |
| self.lm_head = value |
|
|
| |
|
|
| @staticmethod |
| def _causal_mask( |
| attention_mask: torch.Tensor | None, |
| q_len: int, |
| past_len: int, |
| device: torch.device, |
| ) -> tuple[torch.Tensor | None, bool]: |
| """Return the (attn_mask, is_causal) pair to hand to SDPA. |
| |
| The two `None` cases are the fast paths — SDPA can pick a fused kernel |
| only when no explicit mask is materialised: |
| • one query token: it may attend to the entire cache, no mask needed; |
| • uncached prefill: plain `is_causal=True`. |
| Anything else (chunked prefill on top of a cache, or padded batches) |
| needs an explicit mask, because `is_causal=True` aligns to the *top |
| left* of a non-square score matrix and would mask the cache away. |
| """ |
| if attention_mask is not None: |
| if attention_mask.dim() == 4: |
| return attention_mask, False |
| if bool(attention_mask.all()): |
| attention_mask = None |
|
|
| if attention_mask is None: |
| if q_len == 1: |
| return None, False |
| if past_len == 0: |
| return None, True |
|
|
| kv_len = past_len + q_len |
| q_pos = torch.arange(past_len, kv_len, device=device)[:, None] |
| kv_pos = torch.arange(kv_len, device=device)[None, :] |
| mask = (kv_pos <= q_pos)[None, None] |
| if attention_mask is not None: |
| mask = mask & attention_mask[:, None, None, :].bool() |
| |
| |
| mask = mask | ~mask.any(-1, keepdim=True) |
| return mask, False |
|
|
| |
|
|
| def forward( |
| self, |
| input_ids: torch.Tensor | None = None, |
| attention_mask: torch.Tensor | None = None, |
| position_ids: torch.Tensor | None = None, |
| past_key_values: Cache | None = None, |
| labels: torch.Tensor | None = None, |
| use_cache: bool | None = None, |
| logits_to_keep: int | torch.Tensor = 0, |
| **kwargs, |
| ) -> CausalLMOutputWithPast: |
| |
| |
| if use_cache is None: |
| use_cache = self.config.use_cache and labels is None |
| if use_cache and past_key_values is None: |
| past_key_values = DynamicCache() |
| past_len = past_key_values.get_seq_length() if past_key_values is not None else 0 |
|
|
| T = input_ids.shape[1] |
| if position_ids is None: |
| position_ids = torch.arange(past_len, past_len + T, device=input_ids.device) |
| position_ids = position_ids.view(-1, T) |
|
|
| x = self.transformer["wte"](input_ids) |
| cos, sin = self.rotary(position_ids) |
| attn_mask, is_causal = self._causal_mask(attention_mask, T, past_len, x.device) |
|
|
| for block in self.transformer["h"]: |
| x = block(x, cos, sin, attn_mask, is_causal, past_key_values) |
| x = F.rms_norm(x, (x.size(-1),)) |
|
|
| |
| |
| keep = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep |
| logits = self.lm_head(x[:, keep]).float() |
|
|
| loss = None |
| if labels is not None: |
| loss = F.cross_entropy( |
| logits.view(-1, logits.size(-1)), |
| labels.view(-1), |
| ignore_index=-100, |
| ) |
|
|
| return CausalLMOutputWithPast( |
| loss=loss, |
| logits=logits, |
| past_key_values=past_key_values if use_cache else None, |
| ) |
|
|