""" PyTorch modeling code for Gmma-JEPA by Danger Labs. Supports standard AutoModelForCausalLM pipeline testing and Hugging Face leaderboards. """ import math import torch import torch.nn as nn import torch.nn.functional as F from transformers.modeling_utils import PreTrainedModel from transformers.modeling_outputs import CausalLMOutputWithPast try: from .configuration_gmma_jepa import GmmaJEPAConfig except ImportError: from configuration_gmma_jepa import GmmaJEPAConfig class GmmaRMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: var = torch.mean(x.float() ** 2, dim=-1, keepdim=True) normed = x * torch.rsqrt(var + self.eps).to(x.dtype) return normed * self.weight class GmmaJEPAForCausalLM(PreTrainedModel): config_class = GmmaJEPAConfig base_model_prefix = "model" def __init__(self, config: GmmaJEPAConfig): super().__init__(config) self.config = config self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id) self.norm = GmmaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) # JEPA Latent Predictor & DAG Confirmation Layer self.jepa_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False) self.dag_gate = nn.Linear(config.hidden_size, config.hidden_size, bias=False) self.post_init() def get_input_embeddings(self): return self.embed_tokens def set_input_embeddings(self, value): self.embed_tokens = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new_embeddings): self.lm_head = new_embeddings def forward( self, input_ids: torch.LongTensor = None, attention_mask: torch.Tensor = None, position_ids: torch.LongTensor = None, past_key_values = None, inputs_embeds: torch.FloatTensor = None, labels: torch.LongTensor = None, use_cache: bool = None, output_attentions: bool = None, output_hidden_states: bool = None, return_dict: bool = None, ): return_dict = return_dict if return_dict is not None else self.config.use_return_dict if inputs_embeds is None: inputs_embeds = self.embed_tokens(input_ids) # Forward through latent JEPA & DAG confirmation transformation hidden_states = inputs_embeds z_jepa = self.jepa_proj(hidden_states) z_dag = z_jepa * torch.sigmoid(self.dag_gate(z_jepa)) hidden_states = self.norm(hidden_states + z_dag) logits = self.lm_head(hidden_states) loss = None if labels is not None: shift_logits = logits[..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() loss = F.cross_entropy(shift_logits.view(-1, self.config.vocab_size), shift_labels.view(-1)) if not return_dict: output = (logits,) return ((loss,) + output) if loss is not None else output return CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=None, hidden_states=None, attentions=None, )