NewsBERT_1800-1920-Temporal / continuous_time_embedding.py
npedrazzini's picture
Update continuous_time_embedding.py
3eb4aa8 verified
Raw
History Blame Contribute Delete
5.13 kB
"""
Companion module for loading the continuous-time NewsBERT model.
This model is NOT a plain AutoModelForMaskedLM -- it wraps a full
fine-tuned BERT with a continuous sinusoidal (Fourier) time embedding
injected at the input layer. You need this file to load and query it.
Usage:
from continuous_time_embedding import load_continuous_time_model
tokenizer, model = load_continuous_time_model("TextMachineProject/NewsBERT_1800-1920-Temporal")
"""
import os
import math
import torch
import torch.nn as nn
from transformers import AutoTokenizer, AutoModelForMaskedLM
from huggingface_hub import snapshot_download
BASE_MODEL_ID = "TextMachineProject/NewsBERT_1800-1920"
MIN_YEAR = 1800.0
MAX_YEAR = 1920.0
MIN_PERIOD_YEARS = 5.0
N_TIME_FREQS = 24
class ContinuousTimeEmbedding(nn.Module):
"""Sinusoidal (Fourier) features over normalized year, projected to
hidden_size. Nearby years produce nearby embeddings by construction.
Frequency band: lowest = 1 cycle over the full 1800-1920 span, highest =
1 cycle per MIN_PERIOD_YEARS (5 years)."""
def __init__(self, hidden_size, n_freqs=N_TIME_FREQS, min_year=MIN_YEAR, max_year=MAX_YEAR,
min_period_years=MIN_PERIOD_YEARS):
super().__init__()
self.min_year = min_year
self.max_year = max_year
span_years = max_year - min_year
low_freq_per_year = 1.0 / span_years
high_freq_per_year = 1.0 / min_period_years
freqs_per_year = torch.exp(torch.linspace(
math.log(low_freq_per_year), math.log(high_freq_per_year), n_freqs
))
angular_freqs = freqs_per_year * span_years * 2 * math.pi
self.register_buffer("freqs", angular_freqs)
self.proj = nn.Linear(2 * n_freqs, hidden_size)
def forward(self, years: torch.Tensor) -> torch.Tensor:
t = (years - self.min_year) / (self.max_year - self.min_year)
t = t.clamp(0.0, 1.0).unsqueeze(-1)
angles = t * self.freqs
feats = torch.cat([torch.sin(angles), torch.cos(angles)], dim=-1)
return self.proj(feats)
class ContinuousTimeBertForMLM(nn.Module):
"""Full fine-tuned BERT + continuous time embedding, injected into every
token's input embedding, re-normalized via LayerNorm before entering the
transformer stack."""
def __init__(self, model, hidden_size, n_time_freqs=N_TIME_FREQS,
min_year=MIN_YEAR, max_year=MAX_YEAR, inject_mode="all_tokens"):
super().__init__()
self.model = model
self.time_embed = ContinuousTimeEmbedding(hidden_size, n_time_freqs, min_year, max_year)
assert inject_mode in ("all_tokens", "cls_only")
self.inject_mode = inject_mode
self.post_inject_norm = nn.LayerNorm(hidden_size)
def get_input_embeddings_module(self):
return self.model.bert.embeddings
def forward(self, input_ids, attention_mask, years, labels=None):
embeddings_module = self.get_input_embeddings_module()
tok_embeds = embeddings_module(input_ids)
time_vec = self.time_embed(years).unsqueeze(1)
if self.inject_mode == "all_tokens":
tok_embeds = self.post_inject_norm(tok_embeds + time_vec)
else:
tok_embeds = tok_embeds.clone()
tok_embeds[:, 0, :] = self.post_inject_norm(tok_embeds[:, 0, :] + time_vec.squeeze(1))
return self.model(inputs_embeds=tok_embeds, attention_mask=attention_mask, labels=labels)
def save_pretrained(self, save_dir):
os.makedirs(save_dir, exist_ok=True)
self.model.save_pretrained(save_dir)
torch.save(self.time_embed.state_dict(), os.path.join(save_dir, "time_embed.pt"))
torch.save(self.post_inject_norm.state_dict(), os.path.join(save_dir, "post_inject_norm.pt"))
@classmethod
def load_pretrained(cls, save_dir, hidden_size=768, **kwargs):
model = AutoModelForMaskedLM.from_pretrained(save_dir)
obj = cls(model, hidden_size, **kwargs)
obj.time_embed.load_state_dict(torch.load(os.path.join(save_dir, "time_embed.pt"), map_location="cpu"))
obj.post_inject_norm.load_state_dict(torch.load(os.path.join(save_dir, "post_inject_norm.pt"), map_location="cpu"))
return obj
def load_continuous_time_model(repo_id_or_path, device=None, **kwargs):
if os.path.isdir(repo_id_or_path):
local_dir = repo_id_or_path
else:
local_dir = snapshot_download(repo_id_or_path)
try:
candidate_tokenizer = AutoTokenizer.from_pretrained(local_dir)
tokenizer = candidate_tokenizer if len(candidate_tokenizer) >= 1000 else None
except Exception:
tokenizer = None
if tokenizer is None:
print(f"[load_continuous_time_model] No valid tokenizer found in {repo_id_or_path}, "
f"falling back to base model tokenizer: {BASE_MODEL_ID}")
tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL_ID)
model = ContinuousTimeBertForMLM.load_pretrained(local_dir, **kwargs)
device = device or ("cuda" if torch.cuda.is_available() else "cpu")
model.to(device).eval()
return tokenizer, model