""" 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