Download djepa/model/encoder.py from DangerLabs/DM-JEPA: direct link, hf CLI and curl.
- Browser
- Download file 3.27 kB
-
https://huggingface.co/DangerLabs/DM-JEPA/resolve/main/djepa/model/encoder.py
- Command line
-
hf download hf://DangerLabs/DM-JEPA/djepa/model/encoder.py
-
curl -L -o encoder.py https://huggingface.co/DangerLabs/DM-JEPA/resolve/main/djepa/model/encoder.py
3.27 kB
| """ | |
| Encoders for State Context and Candidate Option Criteria. | |
| """ | |
| from typing import Tuple, Optional | |
| import torch | |
| import torch.nn as nn | |
| from transformers import AutoModel, PreTrainedModel | |
| class StateContextEncoder(nn.Module): | |
| """ | |
| Encodes long document state and question instructions into contextual sequence embeddings. | |
| """ | |
| def __init__( | |
| self, | |
| model_name: str = "answerdotai/ModernBERT-base", | |
| pretrained: bool = True, | |
| gradient_checkpointing: bool = True, | |
| ): | |
| super().__init__() | |
| if pretrained: | |
| self.backbone = AutoModel.from_pretrained(model_name) | |
| else: | |
| from transformers import AutoConfig | |
| config = AutoConfig.from_pretrained(model_name) | |
| self.backbone = AutoModel.from_config(config) | |
| if gradient_checkpointing and hasattr(self.backbone, "gradient_checkpointing_enable"): | |
| self.backbone.gradient_checkpointing_enable() | |
| self.d_model = self.backbone.config.hidden_size | |
| def forward( | |
| self, | |
| input_ids: torch.Tensor, | |
| attention_mask: Optional[torch.Tensor] = None, | |
| ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: | |
| """ | |
| Args: | |
| input_ids: [B, L] | |
| attention_mask: [B, L] | |
| Returns: | |
| h_state: [B, L, D] | |
| attention_mask: [B, L] | |
| """ | |
| outputs = self.backbone(input_ids=input_ids, attention_mask=attention_mask) | |
| h_state = outputs.last_hidden_state # [B, L, D] | |
| return h_state, attention_mask | |
| class CriteriaTargetEncoder(nn.Module): | |
| """ | |
| Encodes candidate option criteria into compact semantic target embeddings s_O in latent space. | |
| """ | |
| def __init__(self, backbone: PreTrainedModel, d_model: int): | |
| super().__init__() | |
| self.backbone = backbone | |
| self.d_model = d_model | |
| self.pool_proj = nn.Sequential( | |
| nn.Linear(d_model, d_model), | |
| nn.GELU(), | |
| nn.LayerNorm(d_model), | |
| ) | |
| def forward( | |
| self, | |
| option_input_ids: torch.Tensor, | |
| option_attention_mask: Optional[torch.Tensor] = None, | |
| ) -> torch.Tensor: | |
| """ | |
| Args: | |
| option_input_ids: [B, K, L_opt] | |
| option_attention_mask: [B, K, L_opt] | |
| Returns: | |
| s_options: [B, K, D] target embeddings in latent space | |
| """ | |
| B, K, L = option_input_ids.shape | |
| flat_ids = option_input_ids.view(B * K, L) | |
| flat_mask = ( | |
| option_attention_mask.view(B * K, L) | |
| if option_attention_mask is not None | |
| else None | |
| ) | |
| outputs = self.backbone(input_ids=flat_ids, attention_mask=flat_mask) | |
| hidden = outputs.last_hidden_state # [B*K, L, D] | |
| # Mean pooling over valid tokens | |
| if flat_mask is not None: | |
| mask_expanded = flat_mask.unsqueeze(-1).expand_as(hidden).float() | |
| sum_hidden = torch.sum(hidden * mask_expanded, dim=1) | |
| sum_mask = mask_expanded.sum(dim=1).clamp(min=1e-9) | |
| pooled = sum_hidden / sum_mask | |
| else: | |
| pooled = hidden.mean(dim=1) | |
| s_options = self.pool_proj(pooled) # [B*K, D] | |
| return s_options.view(B, K, self.d_model) | |