File size: 5,129 Bytes
7d7be4a
 
 
ebd9c90
bf5fa66
 
7d7be4a
 
 
 
3eb4aa8
ebd9c90
7d7be4a
 
 
 
 
 
 
 
 
ebd9c90
bf5fa66
7d7be4a
 
 
 
 
 
 
 
 
 
bf5fa66
7d7be4a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a41f6a7
7d7be4a
 
a41f6a7
7d7be4a
 
a41f6a7
7d7be4a
 
 
 
 
 
 
 
 
 
 
 
a41f6a7
 
 
 
 
ebd9c90
bf5fa66
ebd9c90
bf5fa66
 
 
7d7be4a
 
 
 
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
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
"""
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