File size: 6,074 Bytes
3dc9f3b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Missing-feature handling.

For every input feature we store a binary mask (1 = observed, 0 = missing).
Observed values are first standardized per channel (z-score using frozen stats
fit on the training set), so a small-magnitude but informative feature (e.g.
water/binder ratio ~0.4, admixture dosage ~0.01) is not swamped by a large one
(cement content ~500) in the input projection.  Missing entries are then
replaced by a learnable per-channel embedding, observed entries pick up a
learnable per-channel bias, and the value-plus-mask pair is projected through a
small Linear + LayerNorm + SiLU stack.  The mask is fed in both as a
multiplicative gate and as a side channel so the network can still learn whether
a value was measured, imputed, or structurally absent.
"""

from __future__ import annotations

from typing import Optional, Tuple

import torch
from torch import nn


def masked_feature_stats(
    x: torch.Tensor, mask: torch.Tensor, eps: float = 1e-5
) -> Tuple[torch.Tensor, torch.Tensor]:
    """Per-feature (mean, std) over observed entries only (``mask > 0``).

    Columns with no observed entries (or with near-constant values) return
    ``mean=0, std=1`` so standardization is a no-op there.
    """

    mask = (mask > 0).to(x.dtype)
    x = torch.nan_to_num(x, nan=0.0)
    count = mask.sum(dim=0)
    safe = count.clamp(min=1.0)
    mean = (x * mask).sum(dim=0) / safe
    var = (((x - mean) ** 2) * mask).sum(dim=0) / safe
    std = var.clamp(min=0.0).sqrt()
    has = count > 0
    mean = torch.where(has, mean, torch.zeros_like(mean))
    std = torch.where(has & (std > eps), std, torch.ones_like(std))
    return mean, std


def masked_feature_stats_by_type(
    x: torch.Tensor,
    mask: torch.Tensor,
    type_index: torch.Tensor,
    num_types: int,
    eps: float = 1e-5,
) -> Tuple[torch.Tensor, torch.Tensor]:
    """Per-(type, feature) masked stats, shape ``(num_types, D)`` each.

    Used where one encoder multiplexes several feature schemas into the same
    columns (e.g. aggregate vs mortar nodes), so each row type is standardized
    against its own statistics. Types with no rows fall back to ``mean=0, std=1``.
    """

    in_dim = x.size(1)
    means = torch.zeros(num_types, in_dim, dtype=x.dtype)
    stds = torch.ones(num_types, in_dim, dtype=x.dtype)
    for t in range(num_types):
        sel = type_index == t
        if bool(sel.any()):
            means[t], stds[t] = masked_feature_stats(x[sel], mask[sel], eps)
    return means, stds


class MissingFeatureEncoder(nn.Module):
    """Encode a possibly-missing feature vector into a dense hidden representation.

    Parameters
    ----------
    in_dim: int
        Number of raw input channels.
    out_dim: int
        Output hidden dimension.
    """

    def __init__(self, in_dim: int, out_dim: int, num_stat_groups: int = 1):
        super().__init__()
        self.in_dim = in_dim
        self.out_dim = out_dim
        self.num_stat_groups = num_stat_groups

        self.missing_embedding = nn.Parameter(torch.zeros(in_dim))
        nn.init.normal_(self.missing_embedding, std=0.02)
        self.observed_bias = nn.Parameter(torch.zeros(in_dim))

        # Frozen per-feature input standardization (z-score). Default identity;
        # call ``set_feature_stats`` with training-set stats to activate. Stored
        # as buffers (shape ``(num_stat_groups, in_dim)``) so they travel with
        # .to(device) and persist in checkpoints. ``num_stat_groups`` > 1 lets one
        # encoder hold per-node-type / per-edge-type stats, selected per row via a
        # ``type_index`` in ``forward`` (the columns mean different features for
        # different types, so a single shared z-score would blend them).
        self.register_buffer("feat_mean", torch.zeros(num_stat_groups, in_dim))
        self.register_buffer("feat_std", torch.ones(num_stat_groups, in_dim))

        self.proj = nn.Sequential(
            nn.Linear(in_dim * 2, out_dim),
            nn.LayerNorm(out_dim),
            nn.SiLU(),
        )

    @torch.no_grad()
    def set_feature_stats(
        self, mean: torch.Tensor, std: torch.Tensor, eps: float = 1e-5
    ) -> None:
        """Freeze standardization stats (from the train split).

        ``mean`` / ``std`` are ``(in_dim,)`` for a single group or
        ``(num_stat_groups, in_dim)`` for per-type stats.
        """

        mean = torch.as_tensor(mean, dtype=self.feat_mean.dtype, device=self.feat_mean.device)
        std = torch.as_tensor(std, dtype=self.feat_std.dtype, device=self.feat_std.device)
        if mean.dim() == 1:
            mean = mean.unsqueeze(0)
        if std.dim() == 1:
            std = std.unsqueeze(0)
        std = torch.where(std > eps, std, torch.ones_like(std))
        self.feat_mean.copy_(mean.expand_as(self.feat_mean))
        self.feat_std.copy_(std.expand_as(self.feat_std))

    def forward(
        self,
        x: torch.Tensor,
        mask: Optional[torch.Tensor] = None,
        type_index: Optional[torch.Tensor] = None,
    ) -> torch.Tensor:
        if mask is None:
            mask = torch.ones_like(x)
        x = torch.nan_to_num(x, nan=0.0)
        if type_index is None:
            mean, std = self.feat_mean[0], self.feat_std[0]
        else:
            mean = self.feat_mean.index_select(0, type_index)
            std = self.feat_std.index_select(0, type_index)
        x = (x - mean) / std
        imputed = x * mask + self.missing_embedding * (1.0 - mask)
        imputed = imputed + self.observed_bias * mask
        joined = torch.cat([imputed, mask], dim=-1)
        return self.proj(joined)


def apply_random_missingness(
    x: torch.Tensor,
    rate: float = 0.15,
    generator: Optional[torch.Generator] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
    """Randomly drop entries of ``x`` and return ``(x_masked, mask)``."""

    if generator is None:
        mask = (torch.rand_like(x) > rate).float()
    else:
        rnd = torch.rand(x.shape, generator=generator, device=x.device)
        mask = (rnd > rate).float()
    return x * mask, mask