""" PyTorch modeling code for Gmma-JEPA by Danger Labs. Supports standard AutoModelForCausalLM pipeline testing, Hugging Face leaderboards, and native dynamic JEPA LoRA adapter injection. """ import os 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 from transformers.generation import GenerationMixin from safetensors.torch import load_file 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 JEPALoRAAdapterModule(nn.Module): """ Modular Low-Rank Adapter (LoRA) for specialized JEPA domains. """ def __init__(self, in_features: int, out_features: int, r: int = 32, alpha: float = 64.0): super().__init__() self.r = r self.scaling = alpha / r self.lora_A = nn.Linear(in_features, r, bias=False) self.lora_B = nn.Linear(r, out_features, bias=False) nn.init.kaiming_uniform_(self.lora_A.weight, a=2.236) nn.init.zeros_(self.lora_B.weight) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.lora_B(self.lora_A(x)) * self.scaling class GmmaJEPAForCausalLM(PreTrainedModel, GenerationMixin): config_class = GmmaJEPAConfig base_model_prefix = "model" _supports_flash_attn_2 = False _supports_sdpa = True main_input_name = "input_ids" 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) # Dynamic JEPA LoRA Adapters self.jepa_adapters = nn.ModuleDict() self.active_jepa_adapters = [] 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 load_jepa_loras(self, adapter_path: str = None): """ Dynamically loads and registers the 5 JEPA domain LoRAs. """ candidate_paths = [ adapter_path, os.path.join(os.path.dirname(os.path.abspath(__file__)), "jepa_lora_adapters", "adapter_model.safetensors"), "./dist/gmma-jepa-dangerlabs-v1.0/jepa_lora_adapters/adapter_model.safetensors", "./jepa_lora_adapters/adapter_model.safetensors" ] resolved_path = None for p in candidate_paths: if p and os.path.exists(p): resolved_path = p break if resolved_path: weights = load_file(resolved_path) domains = set() for k in weights.keys(): parts = k.split(".") if len(parts) >= 4 and parts[2] == "jepa_adapters": domains.add(parts[3]) for domain in domains: adapter = JEPALoRAAdapterModule(self.config.hidden_size, self.config.hidden_size, r=32, alpha=64.0) adapter = adapter.to(dtype=self.dtype, device=self.device) a_key = f"base_model.model.jepa_adapters.{domain}.lora_A.weight" b_key = f"base_model.model.jepa_adapters.{domain}.lora_B.weight" if a_key in weights and b_key in weights: adapter.lora_A.weight.data.copy_(weights[a_key].to(self.dtype)) adapter.lora_B.weight.data.copy_(weights[b_key].to(self.dtype)) self.jepa_adapters[domain] = adapter if domain not in self.active_jepa_adapters: self.active_jepa_adapters.append(domain) print(f"[gmma-jepa] Successfully attached {len(domains)} JEPA LoRA Adapters: {list(domains)}") else: print(f"[gmma-jepa] Warning: Could not locate adapter_model.safetensors in candidate paths.") def enable_jepa_lora(self, domain_name: str): if domain_name in self.jepa_adapters and domain_name not in self.active_jepa_adapters: self.active_jepa_adapters.append(domain_name) def disable_jepa_lora(self, domain_name: str): if domain_name in self.active_jepa_adapters: self.active_jepa_adapters.remove(domain_name) 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) # Apply active JEPA LoRA Adapters if self.active_jepa_adapters: for domain in self.active_jepa_adapters: if domain in self.jepa_adapters: z_jepa = z_jepa + self.jepa_adapters[domain](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=past_key_values, hidden_states=hidden_states if output_hidden_states else None, attentions=None, ) def prepare_inputs_for_generation( self, input_ids, past_key_values=None, attention_mask=None, inputs_embeds=None, **kwargs ): model_inputs = {"input_ids": input_ids} if past_key_values is not None: model_inputs["past_key_values"] = past_key_values if attention_mask is not None: model_inputs["attention_mask"] = attention_mask if inputs_embeds is not None: model_inputs["inputs_embeds"] = inputs_embeds return model_inputs