File size: 2,362 Bytes
712daaa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
"""
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