Download modeling_dm_jepa.py from DangerLabs/DM-JEPA: direct link, hf CLI and curl.
- Browser
- Download file 2.36 kB
-
https://huggingface.co/DangerLabs/DM-JEPA/resolve/main/modeling_dm_jepa.py
- Command line
-
hf download hf://DangerLabs/DM-JEPA/modeling_dm_jepa.py
-
curl -L -o modeling_dm_jepa.py https://huggingface.co/DangerLabs/DM-JEPA/resolve/main/modeling_dm_jepa.py
2.36 kB
| """ | |
| DM-JEPA: Decision-Making Joint Embedding Predictive Architecture | |
| Published by Danger Labs (https://huggingface.co/DangerLabs) | |
| Non-autoregressive System 1 decision engine operating directly in latent thought space. | |
| """ | |
| import os | |
| import json | |
| import torch | |
| import torch.nn as nn | |
| from typing import Dict, Any, Optional | |
| from djepa.model.djepa import DJEPA | |
| from djepa.dataset.formatter import JevFormatter | |
| from transformers import AutoTokenizer | |
| class DMJEPA(nn.Module): | |
| """ | |
| DM-JEPA by Danger Labs. | |
| High-throughput, ultra-low latency System 1 Decision Model. | |
| """ | |
| def __init__(self, **kwargs): | |
| super().__init__() | |
| self.model = DJEPA(**kwargs) | |
| self.formatter = JevFormatter() | |
| def forward(self, *args, **kwargs): | |
| return self.model(*args, **kwargs) | |
| def from_pretrained(cls, pretrained_model_name_or_path: str, device: Optional[str] = None, **kwargs): | |
| from huggingface_hub import hf_hub_download | |
| from safetensors.torch import load_file | |
| device = device or ("cuda" if torch.cuda.is_available() else "cpu") | |
| path = pretrained_model_name_or_path | |
| if os.path.isdir(path): | |
| config_file = os.path.join(path, "config.json") | |
| safetensors_path = os.path.join(path, "model.safetensors") | |
| weights_file = safetensors_path if os.path.exists(safetensors_path) else os.path.join(path, "pytorch_model.bin") | |
| else: | |
| config_file = hf_hub_download(repo_id=path, filename="config.json") | |
| try: | |
| weights_file = hf_hub_download(repo_id=path, filename="model.safetensors") | |
| except Exception: | |
| weights_file = hf_hub_download(repo_id=path, filename="pytorch_model.bin") | |
| with open(config_file, "r") as f: | |
| cfg = json.load(f) | |
| inst = cls( | |
| model_name=cfg.get("backbone_model", "answerdotai/ModernBERT-base"), | |
| pretrained=False, | |
| ) | |
| if weights_file.endswith(".safetensors"): | |
| state_dict = load_file(weights_file) | |
| else: | |
| loaded = torch.load(weights_file, map_location="cpu") | |
| state_dict = loaded["model_state_dict"] if "model_state_dict" in loaded else loaded | |
| inst.model.load_state_dict(state_dict) | |
| inst.to(device) | |
| inst.eval() | |
| return inst | |