modernnews-event-stream / load_model.py
sirus's picture
Release ModernNews event-stream checkpoint
afe6bad verified
Raw
History Blame Contribute Delete
1.66 kB
"""Download and load the ModernNews event-stream checkpoint."""
from __future__ import annotations
import json
import sys
from pathlib import Path
import torch
from huggingface_hub import snapshot_download
REPO_ID = "sirus/modernnews-event-stream"
def load_model(repo_id: str = REPO_ID, device: str = "cpu"):
repo_dir = Path(
snapshot_download(
repo_id,
allow_patterns=["config.json", "vocabs.json", "event_model.pt", "modernnews/**"],
)
)
sys.path.insert(0, str(repo_dir))
from modernnews.event_model import EventStreamModel
config = json.loads((repo_dir / "config.json").read_text(encoding="utf-8"))
vocabs = json.loads((repo_dir / "vocabs.json").read_text(encoding="utf-8"))
model = EventStreamModel(
num_event_types=len(vocabs["event_types"]),
num_topics=len(vocabs["topics"]),
num_locations=len(vocabs["locations"]),
hidden_size=config["hidden_size"],
tcn_depth=config["tcn_depth"],
tcn_kernel_size=config["tcn_kernel_size"],
dropout=config["dropout"],
)
state_dict = torch.load(repo_dir / config["weights_file"], map_location="cpu", weights_only=True)
model.load_state_dict(state_dict)
return model.to(device).eval(), vocabs, config
if __name__ == "__main__":
loaded_model, loaded_vocabs, loaded_config = load_model()
parameter_count = sum(parameter.numel() for parameter in loaded_model.parameters())
print(f"Loaded {parameter_count:,} parameters")
print({name: len(values) for name, values in loaded_vocabs.items()})
print(f"Expected window size: {loaded_config['window_size']}")