""" LunarisGuardModel — Hub-loadable version. This module is what gets loaded when a user runs: AutoModel.from_pretrained("auren-research/lunaris-guard", trust_remote_code=True) The class is the same dual-head classifier used in training, but configured to load weights from a HF Hub repo via PretrainedConfig + auto_map. Returns a dict with keys: - injection_logits: [B, 2] - safety_logits: [B, 2] - pooled_output: [B, hidden_size] (debug / probing) """ from typing import Optional import torch import torch.nn as nn from transformers import AutoModel, PreTrainedModel from .configuration_lunaris_guard import LunarisGuardConfig class LunarisGuardModel(PreTrainedModel): """Dual-head classifier: injection + content safety, on a ModernBERT backbone.""" config_class = LunarisGuardConfig base_model_prefix = "backbone" supports_gradient_checkpointing = True def __init__(self, config: LunarisGuardConfig): super().__init__(config) self.config = config # Backbone: ModernBERT-base. Loaded fresh; the saved state_dict will # overwrite these weights when from_pretrained() runs. self.backbone = AutoModel.from_pretrained( config.base_model_name, trust_remote_code=True, ) self.dropout = nn.Dropout(config.classifier_dropout) self.injection_head = nn.Linear( config.hidden_size, config.num_injection_classes ) self.safety_head = nn.Linear( config.hidden_size, config.num_safety_classes ) def forward( self, input_ids: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, injection_labels: Optional[torch.Tensor] = None, safety_labels: Optional[torch.Tensor] = None, **kwargs, ): outputs = self.backbone( input_ids=input_ids, attention_mask=attention_mask, return_dict=True, ) # CLS pooling pooled = outputs.last_hidden_state[:, 0, :] pooled = self.dropout(pooled) injection_logits = self.injection_head(pooled) safety_logits = self.safety_head(pooled) return { "injection_logits": injection_logits, "safety_logits": safety_logits, "pooled_output": pooled, }