""" 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) @classmethod 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