Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
172 changes: 151 additions & 21 deletions src/pome/gnn_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
from sklearn.model_selection import KFold
from sklearn.metrics import roc_auc_score
import joblib
from pome.models import GraphAutoencoder
from pome.models import GraphAutoencoder, ValueRegressor
from pome.utils import compute_roc, repeat_pad_to_max_cols, bin_column_non_linear, bin_column_with_na_adjusted, signed_power_bins, get_zscore_bins

def make_deterministic(seed=42):
Expand Down Expand Up @@ -44,6 +44,9 @@ def __init__(self,
discretization_type : str = "z",
enable_imputation : bool = False,
epoch_checkpoints : int = -1,
regressor_epochs : int = 600,
regressor_lr : float = None,
regressor_weight_decay : float = 1e-4,
inductive : bool = False,
cv_folds : int = 3,
cv_eval_every : int = 10,
Expand All @@ -65,6 +68,13 @@ def __init__(self,
self.discretization_type = discretization_type
self.enable_imputation = enable_imputation
self.save_epochs = epoch_checkpoints
# Numeric imputation trains a small regression head on the frozen embeddings: for a
# missing continuous cell it predicts the value from concat(sample_encoding,
# variable_embedding). These knobs tune that head (training budget, learning rate,
# weight decay); regressor_lr=None falls back to the embedder's own lr.
self.regressor_epochs = regressor_epochs
self.regressor_lr = regressor_lr
self.regressor_weight_decay = regressor_weight_decay
# Inductive epoch tuning: when enabled, fit() first runs a sample-holdout CV to pick the
# number of epochs that best generalizes to unseen samples (avoids overfitting the graph),
# then trains on all samples for that many epochs.
Expand Down Expand Up @@ -137,6 +147,8 @@ def fit(self, X, y=None):
self._variable_embeddings = variable_embeddings
self._bin_embeddings = bin_embeddings

if self.enable_imputation and self._cont_vars:
self._fit_value_regressor()
# Restore the epoch cap so the estimator stays re-fittable (inductive tuning temporarily
# lowered self.epochs to the CV-selected value for the training run above).
self.epochs = cap
Expand Down Expand Up @@ -211,9 +223,21 @@ def impute_sample(self, sample_colum : str, na_value : float):
# Extract all positions (i.e. variables) in given sample that possess this missing value.
impute_variables = self._discretized_data.index[self._discretized_data[sample_colum] == na_value].tolist()

# Compute embedding similarities between sample and all to-impute variable nodes.
for variable in impute_variables:
# Retrieve for each sample & category node the respective embedding.
# Continuous variables: predict the value with the regression head trained on the
# frozen embeddings. A continuous variable with no training observations has no head
# entry -> fall back to its observed mean (NaN if it was never observed at all).
if variable in self._cont_vars:
if (hasattr(self, "_value_regressor")
and variable in self._reg_target_stats):
imputed_df.loc[variable, sample_colum] = self._regress_value(
variable, sample_colum)
else:
imputed_df.loc[variable, sample_colum] = self._observed_mean(variable)
continue

# Categorical variables: pick the category whose value node the decoder scores most
# similar to the sample node.
potential_categories = [value for value, is_not_na in self._var_value_dict[variable] if is_not_na]
category_indices = torch.tensor(
[self._value_node_dict[value] for value in potential_categories],
Expand All @@ -233,29 +257,22 @@ def impute_sample(self, sample_colum : str, na_value : float):
with torch.no_grad():
similarities = self._fitted_decoder(self._all_embeddings, edge_index)

# Find category with highest similarity.
# Find category with highest similarity and extract its value.
best_idx = torch.argmax(similarities).item()
closest_value = potential_categories[best_idx]
# Extract value from category name.
category_value = float(closest_value.split("=")[1])

if variable in self._cat_vars:
# Directly set imputed value at corresponding position in dataframe.
imputed_df.loc[variable, sample_colum] = category_value

elif variable in self._cont_vars:
bin_patients = self._discretized_data.columns[
self._discretized_data.loc[variable] == category_value
].tolist()
if bin_patients:
imputed_value = self._X.loc[variable, bin_patients].mean()
else:
imputed_value = np.nan # or global mean or other fallback

imputed_df.loc[variable, sample_colum] = imputed_value
imputed_df.loc[variable, sample_colum] = float(closest_value.split("=")[1])

return imputed_df

def _observed_mean(self, variable):
"""Observed (non-missing) mean of a continuous variable; NaN if never observed.

Fallback for continuous variables that carry no training signal for the regression
head (e.g. a variable missing in every training sample)."""
obs = self._X.loc[variable].astype(float)[
self._observed_continuous_mask(self._X.loc[variable])]
return float(obs.mean()) if len(obs) else np.nan


def _remap_to_populated_bin(self, var, bin_id):
"""Return the nearest bin index to bin_id that has a node in the training graph."""
Expand Down Expand Up @@ -426,6 +443,119 @@ def transform(self, X):
embedding_cols = [f'dim_{i}' for i in range(self.embedding_dimension)]
return pd.DataFrame(new_latent.numpy(), index=new_samples, columns=embedding_cols)

def _observed_continuous_mask(self, series):
"""Boolean mask over a variable's row selecting observed (non-missing) continuous cells."""
col = series.astype(float)
return ~((col == self.non_informative_na) | col.isin(self.informative_nas) | col.isna())

def _reg_target_standardization(self, var):
"""(mean, std) used to standardize a continuous variable's regression targets.

Reuses the z-score bin stats when available; otherwise (nonlinear mode) computes them
from the observed values. std is guarded to 1.0 for constant / degenerate variables.
"""
stats = self._bin_stats[var]
if stats['type'] == 'z':
mean, std = float(stats['mean']), float(stats['std'])
else:
obs = self._X.loc[var].astype(float)[self._observed_continuous_mask(self._X.loc[var])]
mean = float(obs.mean()) if len(obs) else 0.0
std = float(obs.std(ddof=0)) if len(obs) else 1.0
if std == 0 or not np.isfinite(std):
std = 1.0
return mean, std

def _fit_value_regressor(self):
"""Train the regression head on the frozen embeddings for numeric imputation.

The head maps ``concat(sample_encoding, variable_embedding) -> standardized value``,
where the sample encoding is the raw learned node embedding (component_1 + component_2)
from the embedding matrix, taken before the GAT encoder. Encodings are
variable-independent, so a single cached matrix (``_sample_encodings``) serves both
training and inference and the features are identically distributed at both times.
Training is full-batch with an L1 loss and is deterministic under the global seed.
"""
device = next(self.model.parameters()).device
self.model.eval()

# Populated only for variables that actually yield training pairs; a continuous
# variable never observed in training stays out and falls back to its observed mean.
self._reg_target_stats = {}

sample_names = list(self._sample_node_dict.keys())
self._reg_sample_order = {s: i for i, s in enumerate(sample_names)}
sample_node_idx = torch.tensor(
[self._sample_node_dict[s] for s in sample_names],
dtype=torch.long, device=device)

with torch.no_grad():
c1 = self.model.node_embeddings(self.model.node_to_embeddings[:, 0].to(device))
c2 = self.model.node_embeddings(self.model.node_to_embeddings[:, 1].to(device))
training_embeds = c1 + c2

# Pre-encoder sample encodings are variable-independent; cache once for inference.
self._sample_encodings = training_embeds[sample_node_idx].detach().to('cpu')

sample_feats, var_feats, targets = [], [], []
for var in sorted(self._cont_vars):
# A variable with no value nodes carries no learnable signal.
if not self._var_value_dict[var]:
continue

row = self._X.loc[var]
obs_samples = [s for s in row.index[self._observed_continuous_mask(row)]
if s in self._sample_node_dict]
if not obs_samples:
continue

mean, std = self._reg_target_standardization(var)
self._reg_target_stats[var] = (mean, std)
s_idx = [self._sample_node_dict[s] for s in obs_samples]
var_emb = self._variable_embeddings[self._variable_names.index(var)].to(device)

s_idx_t = torch.tensor(s_idx, dtype=torch.long, device=device)
feats = training_embeds[s_idx_t]
tgt = torch.tensor(
[(float(row[s]) - mean) / std for s in obs_samples],
dtype=feats.dtype, device=device)
sample_feats.append(feats)
var_feats.append(var_emb.unsqueeze(0).expand(len(s_idx), -1))
targets.append(tgt)

if not sample_feats:
return # no observed continuous values to learn from

sample_feats = torch.cat(sample_feats, dim=0)
var_feats = torch.cat(var_feats, dim=0)
targets = torch.cat(targets, dim=0)

reg = ValueRegressor(self.embedding_dimension).to(device)
lr = self.regressor_lr if self.regressor_lr is not None else self.lr
optimizer = torch.optim.Adam(reg.parameters(), lr=lr,
weight_decay=self.regressor_weight_decay)

# L1 (not MSE): the imputation metric is MAE, whose minimizer is the conditional
# median -- targeted directly and robust to the heavy-tailed values in these datasets.
reg.train()
for _ in range(self.regressor_epochs):
optimizer.zero_grad()
F.l1_loss(reg(sample_feats, var_feats), targets).backward()
optimizer.step()
reg.eval()
self._value_regressor = reg.to('cpu')

def _regress_value(self, variable, sample):
"""Predict a continuous value for (variable, sample) from the frozen embeddings.

Uses the cached pre-encoder sample encoding and the variable's embedding, then
de-standardizes the head's output with the variable's stored (mean, std).
"""
sample_emb = self._sample_encodings[self._reg_sample_order[sample]].unsqueeze(0)
var_emb = self._variable_embeddings[self._variable_names.index(variable)].unsqueeze(0)
with torch.no_grad():
std_pred = self._value_regressor(sample_emb, var_emb).item()
mean, std = self._reg_target_stats[variable]
return std_pred * std + mean
def _inductive_forward(self, pairs_per_sample):
"""Inject new sample nodes into the frozen training graph and encode all nodes in one pass.

Expand Down
28 changes: 28 additions & 0 deletions src/pome/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,34 @@ def forward(self, node_embeddings, edge_index):
x = self.layer2(x)
return torch.sigmoid(x).squeeze()

class ValueRegressor(torch.nn.Module):
"""Predicts a (standardized) continuous value from a sample and a variable embedding.

Used for regression-based numeric imputation: the concatenation of a sample's latent
embedding and the target variable's embedding is mapped to a single scalar. Trained
after GNN training on frozen embeddings, so it lives outside GraphAutoencoder and is
pickled independently (like the fitted decoder).
"""
def __init__(self, embedding_dim, hidden_dim=None, dropout=0.1):
super().__init__()
# Wider (2x) + deeper than a single hidden layer: one shared head must model the
# (sample, variable) -> value map for every continuous variable, so give it more
# capacity, with dropout for regularisation on the frozen-embedding features.
hidden = hidden_dim or embedding_dim * 2
self.net = nn.Sequential(
nn.Linear(embedding_dim * 2, hidden),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(hidden, hidden),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(hidden, 1),
)

def forward(self, sample_emb, var_emb):
# sample_emb, var_emb: [N, embedding_dim] -> [N]
return self.net(torch.cat([sample_emb, var_emb], dim=1)).squeeze(-1)

# Graph Autoencoder
class GraphAutoencoder(nn.Module):
def __init__(self, num_embeddings, embedding_dim, encoder_layer, node_to_embeddings_index,
Expand Down
73 changes: 73 additions & 0 deletions tests/test_gnn_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,16 @@
import os
import numpy as np
import pytest
import joblib

example_df = pd.read_csv("example.csv", index_col=0)
NA_ENCODING = -99.0
DIMENSION = 16
DEVICE = "cpu"
NUM_SAMPLES = len(example_df.columns)-1
NUM_VARS = len(example_df)
CONT_VARS = [v for v in example_df.index if example_df.loc[v, "type"] in ("numerical", "cont")]
SAMPLE_COLS = [c for c in example_df.columns if c != "type"]

def test_embedder_init():

Expand Down Expand Up @@ -333,6 +336,76 @@ def test_informative_na_z():
embedder.fit(inf_df)
assert embedder._ap >= 0

def _regression_embedder(discretization_type="z", epochs=50):
return Embedder(epochs=epochs,
na_encoding=NA_ENCODING,
embedding_dimension=DIMENSION,
device=DEVICE,
enable_imputation=True,
discretization_type=discretization_type)

def test_regression_imputation_z():
embedder = _regression_embedder("z")
embedder.fit(example_df)
imputed_df = embedder.impute_all(na_value=NA_ENCODING)

assert imputed_df.shape == (NUM_VARS, NUM_SAMPLES)
assert (imputed_df == NA_ENCODING).sum().sum() == 0
assert isinstance(embedder._value_regressor, torch.nn.Module)

def test_regression_imputation_nonlinear():
embedder = _regression_embedder("nonlinear")
embedder.fit(example_df)
imputed_df = embedder.impute_all(na_value=NA_ENCODING)

assert (imputed_df == NA_ENCODING).sum().sum() == 0
# Every continuous variable observed in training gets standardization stats.
assert all(v in embedder._reg_target_stats for v in CONT_VARS)

def test_imputation_fits_regressor_by_default():
# Enabling imputation with continuous variables always trains the regression head.
embedder = Embedder(epochs=50,
na_encoding=NA_ENCODING,
embedding_dimension=DIMENSION,
device=DEVICE,
enable_imputation=True)
embedder.fit(example_df)
assert isinstance(embedder._value_regressor, torch.nn.Module)
imputed_df = embedder.impute_all(na_value=NA_ENCODING)
assert (imputed_df == NA_ENCODING).sum().sum() == 0

def test_regression_output_finite():
embedder = _regression_embedder("z")
embedder.fit(example_df)
imputed_df = embedder.impute_all(na_value=NA_ENCODING)
cont_cells = imputed_df.loc[CONT_VARS, SAMPLE_COLS].to_numpy(dtype=float)
assert np.isfinite(cont_cells).all()

def test_regression_without_imputation():
embedder = Embedder(epochs=50,
na_encoding=NA_ENCODING,
embedding_dimension=DIMENSION,
device=DEVICE,
enable_imputation=False)
embedder.fit(example_df)
assert not hasattr(embedder, "_value_regressor")
with pytest.raises(ValueError):
embedder.impute_all(na_value=NA_ENCODING)

def test_regressor_persistence(tmp_path):
make_deterministic(42)
embedder = _regression_embedder("z")
embedder.fit(example_df)
before = embedder.impute_all(na_value=NA_ENCODING)

path = tmp_path / "reg_embedder.joblib"
joblib.dump(embedder, path)
reloaded = joblib.load(path)
after = reloaded.impute_all(na_value=NA_ENCODING)

assert isinstance(reloaded._value_regressor, torch.nn.Module)
pd.testing.assert_frame_equal(before, after)

def test_informative_na_nonlinear():
# Same with nonlinear discretization covers utils line 133.
INFORMATIVE_NA = -88.0
Expand Down
Loading