From 7895cb4312f1d5fb66409b097b496aedf7bfe7e8 Mon Sep 17 00:00:00 2001 From: Robrecht Cannoodt Date: Tue, 28 Jul 2026 15:28:57 +0200 Subject: [PATCH] fix output dimension in simple_mlp_predict The model predicts mod2, so out_dim must be input_train_mod2.n_vars, not input_test_mod1.n_vars. Same mixup in the GEX2ATAC branch, where ymean (length n_vars(mod2)) is broadcast against n_vars(mod1). --- CHANGELOG.md | 6 ++++++ src/methods/simple_mlp/simple_mlp_predict/script.py | 4 ++-- 2 files changed, 8 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index e9a8396c..2c01bc87 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,3 +1,9 @@ +# task_predict_modality 0.2.0 + +## BUG FIXES + +* `simple_mlp_predict`: Size the model output with `input_train_mod2.n_vars` instead of `input_test_mod1.n_vars`. The two only coincide on the test resources, so this crashed on every real dataset (PR #24). + # task_predict_modality 0.1.1 ## NEW FUNCTIONALITY diff --git a/src/methods/simple_mlp/simple_mlp_predict/script.py b/src/methods/simple_mlp/simple_mlp_predict/script.py index 2509c8ac..aaf31647 100644 --- a/src/methods/simple_mlp/simple_mlp_predict/script.py +++ b/src/methods/simple_mlp/simple_mlp_predict/script.py @@ -61,7 +61,7 @@ def _predict(model,dl): print('Start predict', flush=True) if task == 'GEX2ATAC': - y_pred = ymean*np.ones([input_test_mod1.n_obs, input_test_mod1.n_vars]) + y_pred = ymean*np.ones([input_test_mod1.n_obs, input_train_mod2.n_vars]) else: folds = [0, 1, 2] @@ -82,7 +82,7 @@ def _predict(model,dl): model_inf = MLP.load_from_checkpoint( ckpt, in_dim=X.shape[1], - out_dim=input_test_mod1.n_vars, + out_dim=input_train_mod2.n_vars, ymean=ymean, config=config )