ConcPre / concrete_gnn /missing_features.py
Heterogeneity2025's picture
Upload 29 files
3dc9f3b verified
Raw
History Blame Contribute Delete
6.07 kB
"""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