File size: 12,801 Bytes
ad424e4
31604c9
ad424e4
 
 
 
 
 
 
 
 
 
31604c9
b233cf7
 
ad424e4
 
 
31604c9
 
b233cf7
 
ad424e4
 
 
 
 
 
 
 
 
 
 
 
 
 
b233cf7
ad424e4
 
b233cf7
ad424e4
 
 
 
82ddb20
ad424e4
b233cf7
 
ad424e4
 
 
 
 
 
 
b233cf7
ad424e4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
557f9dd
ad424e4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b233cf7
 
 
ad424e4
 
 
 
 
b233cf7
557f9dd
 
31604c9
557f9dd
31604c9
 
ad424e4
557f9dd
b233cf7
ad424e4
 
 
31604c9
ad424e4
31604c9
 
 
aa64aba
31604c9
ad424e4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
aa64aba
b233cf7
 
31604c9
 
 
 
 
aa64aba
31604c9
 
ad424e4
 
 
 
 
b233cf7
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
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
"""PIMT classification heads — v6 ConcentrationAwarePyramidHead.

The model predicts three static pyramid tiers (top, middle, base notes) using
physics-informed concentration routing. Each tier pools token representations
weighted by OAV-based routing scores, so that top notes are dominated by
high-volatility ingredients and base notes by low-volatility ingredients.

Key design decisions (from v6 spec):
  (a) Temperature-controlled routing softmax with learnable per-tier temperature.
  (b) Physics-as-bias: routing = learned_attention + log_OAV / tau (init at pure physics).
  (c) Scalar conditioning via 2-layer MLP (not raw concatenation).
  (d) Complete padding hygiene with NaN guards.
"""
from __future__ import annotations

import json
import math
from pathlib import Path
from typing import Any

import torch
import torch.nn as nn
import torch.nn.functional as F


def _load_scalar_stats(path: str = "artifacts/scalar_stats_v1.json") -> dict[str, dict[str, float]]:
    """Load training-set scalar statistics for standardization."""
    p = Path(path)
    if p.exists():
        return json.loads(p.read_text())
    # Fallback defaults
    return {
        "log10_oav_sum": {"mean": 35.0, "std": 21.5},
        "x_liquid_sum": {"mean": 1.0, "std": 0.08},
    }


class ConcentrationAwarePyramidHead(nn.Module):
    """Three-tier pyramid head with physics-informed concentration routing.

    Input:
      latent: (B, S, H) — per-token hidden states from the transformer (time already pooled).
      physics: (B, S, 2) — per-token physics states [x_liquid, log10(OAV)].
      src_key_padding_mask: (B, S) — True for padding positions.

    Output: (B, 3, 138) sigmoid probabilities for [top, mid, base].
    """

    def __init__(
        self,
        hidden_dim: int,
        output_dim: int = 138,
        scalar_stats: dict[str, dict[str, float]] | None = None,
        scalar_embed_dim: int = 16,
    ) -> None:
        super().__init__()
        self.hidden_dim = hidden_dim
        self.output_dim = output_dim

        # Load scalar standardization constants
        if scalar_stats is None:
            scalar_stats = _load_scalar_stats()
        oav_stats = scalar_stats.get("log10_oav_sum", {"mean": 35.0, "std": 21.5})
        xliq_stats = scalar_stats.get("x_liquid_sum", {"mean": 1.0, "std": 0.08})
        self.register_buffer("oav_mean", torch.tensor(oav_stats["mean"], dtype=torch.float32))
        self.register_buffer("oav_std", torch.tensor(max(oav_stats["std"], 1e-6), dtype=torch.float32))
        self.register_buffer("xliq_mean", torch.tensor(xliq_stats["mean"], dtype=torch.float32))
        self.register_buffer("xliq_std", torch.tensor(max(xliq_stats["std"], 1e-6), dtype=torch.float32))

        # (a) Learnable per-tier temperature for routing softmax.
        # Init tau=1.0, clamp to [0.1, 10].
        self.log_tau = nn.Parameter(torch.zeros(3))  # log(1.0) = 0

        # (b) Physics-as-bias: learned attention projection (init at zero = pure physics).
        self.attn_proj = nn.Linear(hidden_dim, 1, bias=False)
        nn.init.zeros_(self.attn_proj.weight)  # Start at pure-physics routing

        # (c) Scalar conditioning MLP: [standardized log10(ΣOAV), standardized Σx_liquid] → 16-D.
        self.scalar_mlp = nn.Sequential(
            nn.Linear(2, scalar_embed_dim),
            nn.GELU(),
            nn.Linear(scalar_embed_dim, scalar_embed_dim),
        )

        # Tier-specific linear heads that take the routed pooled vector + scalar embedding.
        pooled_dim = hidden_dim + scalar_embed_dim
        self.top_head = nn.Linear(pooled_dim, output_dim)
        self.mid_head = nn.Linear(pooled_dim, output_dim)
        self.base_head = nn.Linear(pooled_dim, output_dim)

        # Tier-specific scalar stats for routing (which OAV snapshot to use)
        # We use three time-window OAV summaries. For now, they share the overall stats.
        # The routing is per-token: each token gets a routing weight per tier.

    def _compute_routing_weights(
        self,
        latent: torch.Tensor,
        physics: torch.Tensor,
        src_key_padding_mask: torch.Tensor,
    ) -> torch.Tensor:
        """Compute per-tier routing weights for each token.

        Args:
            latent: (B, S, H)
            physics: (B, S, 2) — [x_liquid, log10(OAV)]
            src_key_padding_mask: (B, S) — True for padding

        Returns:
            routing_weights: (B, 3, S) — per-tier softmax weights over tokens.
        """
        B, S, H = latent.shape

        # Extract log10(OAV) per token (physics channel 1)
        log_oav = physics[..., 1]  # (B, S)

        # (b) Learned attention score (init at zero → pure physics at start)
        attn_score = self.attn_proj(latent).squeeze(-1)  # (B, S)

        # Temperature: tau = exp(clamp(log_tau, -2.3, 2.3)) → clamp tau to [0.1, 10]
        tau = torch.exp(torch.clamp(self.log_tau, -2.303, 2.303))  # (3,)

        routing_weights = []
        for tier_idx in range(3):
            # Routing logits = w * attn_score + log_oav / tau
            # w is folded into attn_score (single learned weight per tier)
            # For simplicity, attn_score is shared but tau differs per tier
            logits = attn_score + log_oav / tau[tier_idx]  # (B, S)

            # (d) Mask padding with finfo.min (not hardcoded -1e9)
            neg_mask_val = torch.finfo(logits.dtype).min
            logits = logits.masked_fill(src_key_padding_mask, neg_mask_val)

            # NaN guard: if ALL positions are masked (shouldn't happen), use uniform
            all_masked = src_key_padding_mask.all(dim=1, keepdim=True)  # (B, 1)
            if all_masked.any():
                # Replace fully-masked rows with uniform distribution
                safe_logits = torch.zeros_like(logits)
                safe_logits = safe_logits.masked_fill(src_key_padding_mask, neg_mask_val)
                logits = torch.where(all_masked.expand_as(logits), safe_logits, logits)

            weights = F.softmax(logits, dim=-1)  # (B, S)
            routing_weights.append(weights)

        return torch.stack(routing_weights, dim=1)  # (B, 3, S)

    def _compute_intensity_scalars(
        self,
        physics: torch.Tensor,
        src_key_padding_mask: torch.Tensor,
    ) -> torch.Tensor:
        """Compute standardized intensity scalars: [log10(ΣOAV), Σx_liquid].

        These are computed per-tier (using different OAV snapshots), but for now
        we use the overall per-token physics.

        Args:
            physics: (B, S, 2) — [x_liquid, log10(OAV)]
            src_key_padding_mask: (B, S)

        Returns:
            scalars: (B, 2) — standardized [log10(ΣOAV), Σx_liquid]
        """
        B, S, _ = physics.shape

        # Zero out padding in the sums
        mask = (~src_key_padding_mask).float().unsqueeze(-1)  # (B, S, 1)
        masked_physics = physics * mask  # (B, S, 2)

        # Sum over tokens (axis 1)
        log_oav_sum = masked_physics[..., 1].sum(dim=1)  # (B,)
        x_liq_sum = masked_physics[..., 0].sum(dim=1)  # (B,)

        # Standardize using training-set statistics
        std_oav = (log_oav_sum - self.oav_mean) / self.oav_std  # (B,)
        std_xliq = (x_liq_sum - self.xliq_mean) / self.xliq_std  # (B,)

        return torch.stack([std_oav, std_xliq], dim=-1)  # (B, 2)

    def forward(
        self,
        latent: torch.Tensor,
        physics: torch.Tensor,
        src_key_padding_mask: torch.Tensor,
    ) -> dict[str, torch.Tensor]:
        """Forward pass.

        Args:
            latent: (B, S, H) — per-token hidden states.
            physics: (B, S, 2) — [x_liquid, log10(OAV)] per token.
            src_key_padding_mask: (B, S) — True for padding.

        Returns:
            dict with:
              "pyramid": (B, 3, 138) sigmoid probabilities
              "routing_weights": (B, 3, S) routing weights (for diagnostics)
              "scalar_embedding": (B, 16) scalar conditioning embedding
        """
        B, S, H = latent.shape

        # (a, b) Compute routing weights
        routing = self._compute_routing_weights(latent, physics, src_key_padding_mask)  # (B, 3, S)

        # (c) Compute scalar conditioning
        scalars = self._compute_intensity_scalars(physics, src_key_padding_mask)  # (B, 2)
        scalar_emb = self.scalar_mlp(scalars)  # (B, 16)

        # Routed pooling: weighted sum of token representations per tier
        # routing: (B, 3, S), latent: (B, S, H)
        pooled = torch.bmm(routing, latent)  # (B, 3, H)

        # Concatenate scalar embedding to each tier
        scalar_expanded = scalar_emb.unsqueeze(1).expand(-1, 3, -1)  # (B, 3, 16)
        pooled_with_scalar = torch.cat([pooled, scalar_expanded], dim=-1)  # (B, 3, H+16)

        # Tier-specific linear heads + sigmoid
        top = torch.sigmoid(self.top_head(pooled_with_scalar[:, 0]))  # (B, 138)
        mid = torch.sigmoid(self.mid_head(pooled_with_scalar[:, 1]))
        base = torch.sigmoid(self.base_head(pooled_with_scalar[:, 2]))

        pyramid = torch.stack([top, mid, base], dim=1)  # (B, 3, 138)

        return {
            "pyramid": pyramid,
            "routing_weights": routing,
            "scalar_embedding": scalar_emb,
        }


class PIMTHeads(nn.Module):
    """Container for the concentration-aware pyramid head and subjective heads.

    Accepts the full 4D latent from the transformer (B, T, S, H) and pools
    over time before feeding to the pyramid head.
    """

    def __init__(
        self,
        hidden_dim: int,
        objective_dim: int = 138,
        seasonality_classes: int = 4,
        wearability_classes: int = 2,
        scalar_stats: dict[str, dict[str, float]] | None = None,
    ) -> None:
        super().__init__()
        self.objective_head = ConcentrationAwarePyramidHead(
            hidden_dim, objective_dim, scalar_stats=scalar_stats
        )

        # Subjective heads operate on a globally-pooled representation.
        self.seasonality_head = nn.Linear(hidden_dim, seasonality_classes)
        self.gender_head = nn.Linear(hidden_dim, 1)
        self.wearability_head = nn.Linear(hidden_dim, wearability_classes)
        self.substantivity_head = nn.Linear(hidden_dim, 1)

    def forward(
        self,
        latent: torch.Tensor,
        physics: torch.Tensor | None = None,
        src_key_padding_mask: torch.Tensor | None = None,
    ) -> dict[str, Any]:
        """Forward pass.

        Args:
            latent: (B, T, S, H) from the transformer encoder.
            physics: (B, T, S, 2) physics states. If None, uses zero physics.
            src_key_padding_mask: (B, S) padding mask. If None, no padding.

        Returns:
            dict with objective, subjective, alignment, and diagnostics.
        """
        # Pool over time dimension to get (B, S, H)
        latent_pooled_time = latent.mean(dim=1)  # (B, S, H)

        B, S, H = latent_pooled_time.shape

        # Handle physics input
        if physics is not None:
            # Pool physics over time too: (B, T, S, 2) → (B, S, 2)
            physics_pooled = physics.mean(dim=1)  # (B, S, 2)
        else:
            physics_pooled = latent_pooled_time.new_zeros(B, S, 2)

        # Handle padding mask
        if src_key_padding_mask is None:
            src_key_padding_mask = torch.zeros(B, S, dtype=torch.bool, device=latent.device)

        # Concentration-aware pyramid prediction
        pyramid_result = self.objective_head(latent_pooled_time, physics_pooled, src_key_padding_mask)
        pyramid = pyramid_result["pyramid"]  # (B, 3, 138)

        # Alignment head: mean of tier logits for cosine alignment loss
        pooled_global = latent_pooled_time.mean(dim=1)  # (B, H)
        alignment_raw = pyramid_result["pyramid"].mean(dim=1)  # (B, 138) — already sigmoided
        alignment = alignment_raw  # Use the mean of sigmoid outputs as alignment target

        # Subjective heads
        seasonality = self.seasonality_head(pooled_global)
        gender = self.gender_head(pooled_global)
        wearability = self.wearability_head(pooled_global)
        substantivity = self.substantivity_head(pooled_global)

        return {
            "objective": pyramid,
            "subjective": {
                "seasonality": seasonality,
                "gender_profile": gender,
                "wearability": wearability,
                "substantivity": substantivity,
            },
            "alignment": alignment,
            "diagnostics": {
                "routing_weights": pyramid_result["routing_weights"],
                "scalar_embedding": pyramid_result["scalar_embedding"],
                "tau": torch.exp(torch.clamp(self.objective_head.log_tau, -2.303, 2.303)),
            },
        }