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
|