A unified framework for correcting batch effects and integrating multi-omics data.
The 2 matches
- [1] § Methods › Hyperparameter setting ↔ moDAmix.py, lines 225–246 · score 0.67 · domain discriminators, single omics feature, multi omics feature, Adam, softmax, PyTorch
- [2] § Methods › Phase 1: pre-training feature extractors and classifier ↔ moDAmix.py, lines 225–246 · score 0.62 · cross entropy loss, multi omics feature, softmax, classifier, domain
Paper
Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC
The paper is loaded when this pane is shown.
The authors' code
Python · 681 lines · 27 KB · MIT · 2 matches
- import torch
- from torch import nn
- from torch.utils.data import DataLoader
- from torch.utils.data import Dataset
- import pandas as pd
- import os
- import numpy as np
- import sys
- class SourceDataset(Dataset):
- def __init__(self, x_data, x_gene, y_data):
- self.x_data = x_data
- self.x_gene = x_gene
- self.y_data = y_data
- def __getitem__(self, index):
- return self.x_data[index], self.x_gene[index], self.y_data[index]
- def __len__(self):
- return self.x_data.shape[0]
- class UnlabelDataset(Dataset):
- def __init__(self, x_data, x_gene):
- self.x_data = x_data
- self.x_gene = x_gene
- def __getitem__(self, index):
- return self.x_data[index], self.x_gene[index]
- def __len__(self):
- return self.x_data.shape[0]
- class DomainDataset(Dataset) :
- def __init__(self, x_data, x_gene, y_data, z_data):
- self.x_data = x_data
- self.x_gene = x_gene
- self.y_data = y_data
- self.z_data = z_data
- def __getitem__(self, index):
- return self.x_data[index], self.x_gene[index], self.y_data[index], self.z_data[index]
- def __len__(self):
- return self.x_data.shape[0]
- device = (
- "cuda"
- if torch.cuda.is_available()
- else "cpu"
- )
- if device == "cuda" :
- os.environ["CUDA_VISIBLE_DEVICES"] = "0"
- print(f"Using {device} device")
- result_dir = "./results"
- os.makedirs(result_dir, exist_ok = True)
- data_dir = sys.argv[1]
- sourceDataDir = data_dir
- targetDataDir = data_dir
- x_filename = os.path.join(sourceDataDir, sys.argv[2])
- y_filename = os.path.join(sourceDataDir, sys.argv[4])
- target_filename = os.path.join(targetDataDir, sys.argv[5])
- x_gene_filename = os.path.join(sourceDataDir, sys.argv[3])
- target_gene_filename = os.path.join(targetDataDir, sys.argv[6])
- raw_x = pd.read_csv(x_filename, index_col = 0)
- raw_y = pd.read_csv(y_filename, index_col = 0)
- raw_target_x = pd.read_csv(target_filename, index_col = 0)
- raw_x_gene = pd.read_csv(x_gene_filename, index_col = 0)
- raw_target_x_gene = pd.read_csv(target_gene_filename, index_col = 0)
- sample_id_list = raw_x.index.tolist()
- sample_id_list.extend(raw_target_x.index.tolist())
- sample_id_list_gene = raw_x_gene.index.tolist()
- sample_id_list_gene.extend(raw_target_x_gene.index.tolist())
- raw_target_domain_y = raw_target_x['domain_idx'].tolist()
- raw_target_domain_y_gene = raw_target_x_gene['domain_idx'].tolist()
- raw_y_colname = raw_y.columns.tolist()[0]
- y_train = raw_y[raw_y_colname].tolist()
- num_subtype = len(set(y_train))
- y_train = np.array(y_train)
- del raw_target_x['domain_idx']
- del raw_target_x['Batch']
- del raw_target_x_gene['domain_idx']
- del raw_target_x_gene['Batch']
- raw_target_x = raw_target_x.values
- x_train = raw_x.values
- raw_target_x_gene = raw_target_x_gene.values
- x_train_gene = raw_x_gene.values
- domain_x = np.append(x_train, raw_target_x, axis = 0)
- domain_x_gene = np.append(x_train_gene, raw_target_x_gene, axis = 0)
- raw_source_domain_y = np.zeros(len(y_train), dtype = int) # TCGA label : 0
- domain_y = np.append(raw_source_domain_y, raw_target_domain_y)
- raw_source_domain_y_gene = np.zeros(len(raw_x_gene), dtype = int) # TCGA label : 0
- domain_y_gene = np.append(raw_source_domain_y_gene, raw_target_domain_y_gene)
- num_domain = len(set(domain_y))
- x_train = torch.from_numpy(x_train)
- y_train = torch.from_numpy(y_train)
- domain_x = torch.from_numpy(domain_x)
- domain_y = torch.from_numpy(domain_y)
- x_train_gene = torch.from_numpy(x_train_gene)
- domain_x_gene = torch.from_numpy(domain_x_gene)
- domain_y_gene = torch.from_numpy(domain_y_gene)
- target_x = torch.from_numpy(raw_target_x)
- target_x_gene = torch.from_numpy(raw_target_x_gene)
- target_init_y = torch.randint(low=0, high=num_subtype, size = (len(target_x),))
- #domain_z : domain_subtype
- domain_z = torch.cat((y_train, target_init_y), 0)
- num_feature = len(x_train[0])
- num_feature_gene = len(x_train_gene[0])
- num_train = len(x_train)
- num_test = len(raw_target_x)
- train_dataset = SourceDataset(x_train, x_train_gene, y_train)
- domain_dataset = DomainDataset(domain_x, domain_x_gene, domain_y, domain_z)
- target_dataset = SourceDataset(target_x, target_x_gene, target_init_y)
- batch_size = 128
- target_batch_size = 128
- test_target_batch_size = 64
- train_dataloader = DataLoader(train_dataset, batch_size = batch_size, shuffle = True)
- domain_dataloader = DataLoader(domain_dataset, batch_size = batch_size, shuffle = True)
- target_dataloader = DataLoader(target_dataset, batch_size = target_batch_size, shuffle = False)
- n_fe_embed1 = 1024
- n_fe_embed2 = 512
- n_mo_fe_embed1 = 512
- n_mo_fe_embed2 = 256
- n_c_h1 = 128
- n_c_h2 = 64
- n_d_h1 = 256
- n_d_h2 = 64
- class SingleOmicsFeatureExtractor(nn.Module) :
- def __init__(self, n_input) :
- super().__init__()
- self.feature_layer = nn.Sequential(
- nn.Linear(n_input, n_fe_embed1),
- nn.LeakyReLU(),
- nn.Linear(n_fe_embed1, n_fe_embed2),
- nn.LeakyReLU()
- )
- def forward(self, x) :
- embedding = self.feature_layer(x)
- return embedding
- class MultiOmicsFeatureExtractor(nn.Module) :
- def __init__(self) :
- super().__init__()
- self.feature_layer = nn.Sequential(
- nn.Linear(n_fe_embed2*2, n_mo_fe_embed1),
- nn.LeakyReLU(),
- nn.Linear(n_mo_fe_embed1, n_mo_fe_embed2),
- nn.LeakyReLU()
- )
- def forward(self, x) :
- embedding = self.feature_layer(x)
- return embedding
- class DomainDiscriminator(nn.Module) :
- def __init__(self, n_fe_h2) :
- super().__init__()
- self.disc_layer = nn.Sequential(
- nn.Linear(n_fe_h2, n_d_h1),
- nn.LeakyReLU(),
- nn.Linear(n_d_h1, n_d_h2),
- nn.LeakyReLU(),
- nn.Linear(n_d_h2, num_domain)
- )
- def forward(self, x) :
- domain_logits = self.disc_layer(x)
- return domain_logits
- class SubtypeClassifier(nn.Module):
- def __init__(self):
- super().__init__()
- #self.flatten = nn.Flatten()
- self.linear_relu_stack = nn.Sequential(
- nn.Linear(n_mo_fe_embed2, n_c_h1),
- nn.LeakyReLU(),
- nn.Linear(n_c_h1, n_c_h2),
- nn.LeakyReLU(),
- nn.Linear(n_c_h2, num_subtype)
- )
- def forward(self, x):
- logits = self.linear_relu_stack(x)
- return logits
- fe_model_methyl = SingleOmicsFeatureExtractor(num_feature).to(device)
- fe_model_gene = SingleOmicsFeatureExtractor(num_feature_gene).to(device)
- fe_model_multiomics = MultiOmicsFeatureExtractor().to(device)
- domain_disc_methyl_model = DomainDiscriminator(n_fe_embed2).to(device)
- domain_disc_gene_model = DomainDiscriminator(n_fe_embed2).to(device)
- domain_disc_multiomics_model = DomainDiscriminator(n_mo_fe_embed2).to(device)
- subtype_pred_model = SubtypeClassifier().to(device)
- c_loss = nn.CrossEntropyLoss() # Already have softmax
- domain_loss = nn.CrossEntropyLoss() # Already have softmax
- fe_methyl_optimizer = torch.optim.Adam(fe_model_methyl.parameters(), lr=1e-4)
- fe_gene_optimizer = torch.optim.Adam(fe_model_gene.parameters(), lr=1e-4)
- fe_multiomics_optimizer = torch.optim.Adam(fe_model_multiomics.parameters(), lr=1e-4)
- c_optimizer = torch.optim.Adam(subtype_pred_model.parameters(), lr=1e-5)
- d_methyl_optimizer = torch.optim.Adam(domain_disc_methyl_model.parameters(), lr=1e-6)
- d_gene_optimizer = torch.optim.Adam(domain_disc_gene_model.parameters(), lr=1e-6)
- d_multiomics_optimizer = torch.optim.Adam(domain_disc_multiomics_model.parameters(), lr=1e-6)
- def pretrain_classifier(epoch, dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, c_model, c_loss, fe_methyl_optimizer, fe_gene_optimizer, fe_multiomics_optimizer, c_optimizer):
- size = len(dataloader.dataset)
- correct = 0
- for batch, (X, X_gene, y) in enumerate(dataloader):
- X, X_gene, y = X.to(device), X_gene.to(device), y.to(device)
- X = X.float()
- X_gene = X_gene.float()
- embed_methyl = fe_model_methyl(X)
- embed_gene = fe_model_gene(X_gene)
- embed_concated = torch.cat((embed_methyl, embed_gene), 1)
- embed_multiomics = fe_model_multiomics(embed_concated)
- pred = c_model(embed_multiomics)
- loss = c_loss(pred, y)
- fe_methyl_optimizer.zero_grad()
- fe_gene_optimizer.zero_grad()
- fe_multiomics_optimizer.zero_grad()
- c_optimizer.zero_grad()
- loss.backward()
- fe_methyl_optimizer.step()
- fe_gene_optimizer.step()
- fe_multiomics_optimizer.step()
- c_optimizer.step()
- correct += (pred.argmax(1) == y).type(torch.float).sum().item()
- loss = loss.item()
- correct /= size
- if epoch % 10 == 0 :
- print(f"[PT Epoch {epoch+1}] \tTraining loss: {loss:>5f}, Training Accuracy: {(100*correct):>0.2f}%")
- def adversarial_train_disc_single_omics(epoch, dataloader, fe_model_methyl, fe_model_gene, d_model_methyl, d_model_gene, domain_loss, fe_methyl_optimizer, fe_gene_optimizer, d_methyl_optimizer, d_gene_optimizer) :
- size = len(dataloader.dataset)
- correct_methyl = 0
- correct_gene = 0
- for batch, (X, X_gene, y, z_subtype) in enumerate(dataloader):
- X, X_gene, y = X.to(device), X_gene.to(device), y.to(device)
- X = X.float()
- X_gene = X_gene.float()
- embed_methyl = fe_model_methyl(X)
- pred = d_model_methyl(embed_methyl)
- d_loss = domain_loss(pred, y)
- # Backpropagation for methyl
- fe_methyl_optimizer.zero_grad()
- d_methyl_optimizer.zero_grad()
- d_loss.backward()
- d_methyl_optimizer.step()
- correct_methyl += (pred.argmax(1) == y).type(torch.float).sum().item()
- #
- embed_gene = fe_model_gene(X_gene)
- pred_gene = d_model_gene(embed_gene)
- d_gene_loss = domain_loss(pred_gene, y)
- fe_gene_optimizer.zero_grad()
- d_gene_optimizer.zero_grad()
- d_gene_loss.backward()
- d_gene_optimizer.step()
- correct_gene += (pred_gene.argmax(1) == y).type(torch.float).sum().item()
- d_loss = d_loss.item()
- d_gene_loss = d_gene_loss.item()
- correct_methyl /= size
- correct_gene /= size
- if t % 10 == 0 :
- print(f"[AT-S Epoch {epoch+1}] Disc me loss: {d_loss:>5f} (Acc: {(100*correct_methyl):>0.2f}%), Disc gene loss: {d_gene_loss:>5f} (Acc: {(100*correct_gene):>0.2f}%)", end = ", ")
- def adversarial_train_disc_multiomics(epoch, dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, d_model_multiomics, domain_loss, fe_multiomics_optimizer, d_multiomics_optimizer) :
- size = len(dataloader.dataset)
- correct = 0
- for batch, (X, X_gene, y, z_subtype) in enumerate(dataloader):
- X, X_gene, y = X.to(device), X_gene.to(device), y.to(device)
- X = X.float()
- X_gene = X_gene.float()
- #
- embed_methyl = fe_model_methyl(X)
- embed_gene = fe_model_gene(X_gene)
- embed_concated = torch.cat((embed_methyl, embed_gene), 1)
- embed_multiomics = fe_model_multiomics(embed_concated)
- #
- pred = d_model_multiomics(embed_multiomics)
- d_loss = domain_loss(pred, y)
- # Backpropagation
- fe_multiomics_optimizer.zero_grad()
- d_multiomics_optimizer.zero_grad()
- d_loss.backward()
- d_multiomics_optimizer.step()
- correct += (pred.argmax(1) == y).type(torch.float).sum().item()
- d_loss = d_loss.item()
- correct /= size
- if t % 10 == 0 :
- print(f"[AT-M Epoch {epoch+1}] Disc loss: {d_loss:>5f}, Training Accuracy: {(100*correct):>0.2f}%", end = ", ")
- def adversarial_train_fe_single_omics(epoch, dataloader, fe_model_methyl, fe_model_gene, d_model_methyl, d_model_gene, domain_loss, fe_methyl_optimizer, fe_gene_optimizer, d_methyl_optimizer, d_gene_optimizer) :
- size = len(dataloader.dataset)
- for batch, (X, X_gene, y, z_subtype) in enumerate(dataloader):
- X, X_gene, y = X.to(device), X_gene.to(device), y.to(device)
- X = X.float()
- X_gene = X_gene.float()
- embed_methyl = fe_model_methyl(X)
- pred = d_model_methyl(embed_methyl)
- fake_y = torch.randint(low=0, high=num_domain, size = (len(y),))
- fake_y = fake_y.to(device)
- g_loss = domain_loss(pred, fake_y)
- # Backpropagation
- fe_methyl_optimizer.zero_grad()
- d_methyl_optimizer.zero_grad()
- g_loss.backward()
- fe_methyl_optimizer.step()
- # Gene
- embed_gene = fe_model_gene(X_gene)
- pred_gene = d_model_gene(embed_gene)
- fake_y_gene = torch.randint(low=0, high=num_domain, size = (len(y),))
- fake_y_gene = fake_y_gene.to(device)
- g_gene_loss = domain_loss(pred_gene, fake_y)
- fe_gene_optimizer.zero_grad()
- d_gene_optimizer.zero_grad()
- g_gene_loss.backward()
- fe_gene_optimizer.step()
- g_loss = g_loss.item()
- g_gene_loss = g_gene_loss.item()
- if epoch % 10 == 0:
- print(f"Gen methyl loss: {g_loss:>5f}, Gene gene loss: {g_gene_loss:>5f}")
- def adversarial_train_fe_multiomics(epoch, dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, d_model_multiomics, domain_loss, fe_multiomics_optimizer, d_multiomics_optimizer) :
- size = len(dataloader.dataset)
- for batch, (X, X_gene, y, z_subtype) in enumerate(dataloader):
- X, X_gene, y = X.to(device), X_gene.to(device), y.to(device)
- X = X.float()
- X_gene = X_gene.float()
- #
- embed_methyl = fe_model_methyl(X)
- embed_gene = fe_model_gene(X_gene)
- embed_concated = torch.cat((embed_methyl, embed_gene), 1)
- embed_multiomics = fe_model_multiomics(embed_concated)
- #
- pred = d_model_multiomics(embed_multiomics)
- fake_y = torch.randint(low=0, high=num_domain, size = (len(y),))
- fake_y = fake_y.to(device)
- g_loss = domain_loss(pred, fake_y)
- # Backpropagation
- fe_multiomics_optimizer.zero_grad()
- d_multiomics_optimizer.zero_grad()
- g_loss.backward()
- fe_multiomics_optimizer.step()
- g_loss = g_loss.item()
- if epoch % 10 == 0:
- print(f"Gen multi loss: {g_loss:>5f}")
- def class_alignment_train(epoch, domain_dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, fe_methyl_optimizer, fe_gene_optimizer, fe_multiomics_optimizer, c_optimizer) :
- for batch, (X, X_gene, y_domain, z_subtype) in enumerate(domain_dataloader):
- X, X_gene, y_domain, z_subtype = X.to(device), X_gene.to(device), y_domain.to(device), z_subtype.to(device)
- X = X.float()
- X_gene = X_gene.float()
- batch_subtype_list = z_subtype.unique()
- #
- embed_methyl = fe_model_methyl(X)
- embed_gene = fe_model_gene(X_gene)
- embed_concated = torch.cat((embed_methyl, embed_gene), 1)
- X_embed = fe_model_multiomics(embed_concated)
- #
- align_loss = torch.zeros((1) ,dtype = torch.float64)
- align_loss = align_loss.to(device)
- #
- for subtype in batch_subtype_list :
- sample_idx_list = (z_subtype == subtype).nonzero(as_tuple = True)[0]
- if len(sample_idx_list) < 1 :
- continue
- #else :
- tmp_x = X_embed[sample_idx_list]
- tmp_y = y_domain[sample_idx_list]
- tmp_z = z_subtype[sample_idx_list]
- batch_domain_list = tmp_y.unique()
- domain_centroid_stack = []
- for domain in batch_domain_list :
- domain_idx_list = (tmp_y == domain).nonzero(as_tuple = True)[0]
- if len(domain_idx_list) != 1 :
- tmp_x_domain = tmp_x[domain_idx_list]
- tmp_centroid = torch.div(torch.sum(tmp_x_domain, dim = 0), len(domain_idx_list))
- domain_centroid_stack.append(tmp_centroid)
- if len(domain_centroid_stack) == 0 :
- continue
- else :
- domain_centroid_stack = torch.stack(domain_centroid_stack)
- subtype_centroid = torch.mean(domain_centroid_stack, dim = 0)
- # Duplicate the subtype centroid to get dist with each domain_centroid
- subtype_centroid_stack = []
- for i in range(len(domain_centroid_stack)) :
- subtype_centroid_stack.append(subtype_centroid)
- subtype_centroid_stack = torch.stack(subtype_centroid_stack)
- pdist_stack = nn.L1Loss()(subtype_centroid_stack, domain_centroid_stack)
- align_loss += torch.mean(pdist_stack, dim = 0)
- if align_loss == 0.0 :
- continue
- align_loss = align_loss / len(batch_subtype_list)
- fe_methyl_optimizer.zero_grad()
- fe_gene_optimizer.zero_grad()
- fe_multiomics_optimizer.zero_grad()
- c_optimizer.zero_grad()
- align_loss.backward()
- fe_methyl_optimizer.step()
- fe_gene_optimizer.step()
- fe_multiomics_optimizer.step()
- c_optimizer.step()
- align_loss = align_loss.item()
- if epoch % 10 == 0 :
- print(f"[CA Epoch {epoch+1}] align loss: {align_loss:>5f}\n")
- def ssl_train_classifier(epoch, source_dataloader, target_dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, c_model, c_loss, fe_methyl_optimizer, fe_gene_optimizer, fe_multiomics_optimizer, c_optimizer) :
- source_size = len(source_dataloader.dataset)
- target_size = len(target_dataloader.dataset)
- #
- # 1. Obtain the pseudo-label for target dataset
- #
- target_pseudo_label = torch.empty((0), dtype = torch.int64)
- target_pseudo_label = target_pseudo_label.to(device)
- #
- for batch, (target_X, target_X_gene, target_y) in enumerate(target_dataloader):
- target_X, target_X_gene, target_y = target_X.to(device), target_X_gene.to(device), target_y.to(device)
- target_X = target_X.float()
- target_X_gene = target_X_gene.float()
- #
- embed_methyl = fe_model_methyl(target_X)
- embed_gene = fe_model_gene(target_X_gene)
- embed_concated = torch.cat((embed_methyl, embed_gene), 1)
- extracted_feature = fe_model_multiomics(embed_concated)
- #
- #extracted_feature = fe_model(target_X)
- batch_target_pred = c_model(extracted_feature)
- batch_pseudo_label = batch_target_pred.argmax(1)
- target_pseudo_label = torch.cat((target_pseudo_label, batch_pseudo_label), 0)
- if batch == 0 :
- target_loss = c_loss(batch_target_pred, target_y)
- else :
- target_loss = target_loss + c_loss(batch_target_pred, target_y)
- target_loss = target_loss / (batch + 1)
- #
- # Define alpha value
- alpha_f = 0.01
- t1 = 100
- t2 = 200
- if epoch < t1 :
- alpha = 0
- elif epoch < t2 :
- alpha = (epoch - t1) / (t2 - t1) * alpha_f
- else :
- alpha = alpha_f
- #
- # 2. Calculate the loss for the source dataset
- #
- correct = 0
- for batch, (source_X, source_X_gene, source_y) in enumerate(source_dataloader):
- source_X, source_X_gene, source_y = source_X.to(device), source_X_gene.to(device), source_y.to(device)
- source_X = source_X.float()
- source_X_gene = source_X_gene.float()
- #
- source_embed_methyl = fe_model_methyl(source_X)
- source_embed_gene = fe_model_gene(source_X_gene)
- source_embed_concated = torch.cat((source_embed_methyl, source_embed_gene), 1)
- source_extracted_feature = fe_model_multiomics(source_embed_concated)
- #
- #source_extracted_feature = fe_model(source_X)
- source_pred = c_model(source_extracted_feature)
- source_loss = c_loss(source_pred, source_y)
- ssl_loss = source_loss + alpha * target_loss
- # Backpropogation
- target_loss.detach_()
- fe_methyl_optimizer.zero_grad()
- fe_gene_optimizer.zero_grad()
- fe_multiomics_optimizer.zero_grad()
- c_optimizer.zero_grad()
- ssl_loss.backward() #retain_graph=True
- fe_methyl_optimizer.step()
- fe_gene_optimizer.step()
- fe_multiomics_optimizer.step()
- c_optimizer.step()
- correct += (source_pred.argmax(1) == source_y).type(torch.float).sum().item()
- ssl_loss = ssl_loss.item()
- source_loss = source_loss.item()
- target_loss = target_loss.item()
- correct /= source_size
- if epoch % 10 == 0 :
- print(f"[SSL Epoch {epoch+1}] alpha : {alpha:>3f}, SSL loss: {ssl_loss:>5f}, source loss: {source_loss:>5f}, target loss: {target_loss:>4f}, source ACC: {(100*correct):>0.2f}%\n")
- return target_pseudo_label
- def get_embed(dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, c_model) :
- fe_model_methyl.eval()
- fe_model_gene.eval()
- fe_model_multiomics.eval()
- c_model.eval()
- X_embed_list = []
- y_list = []
- with torch.no_grad() :
- for batch, (X, X_gene, y) in enumerate(dataloader):
- X, X_gene, y = X.to(device), X_gene.to(device), y.to(device)
- X = X.float()
- #
- embed_methyl = fe_model_methyl(X)
- embed_gene = fe_model_gene(X_gene)
- embed_concated = torch.cat((embed_methyl, embed_gene), 1)
- X_embed = fe_model_multiomics(embed_concated)
- #
- X_embed_list.append(X_embed)
- y_list.append(y)
- X_embed_list = torch.cat(X_embed_list, 0)
- y_list = torch.cat(y_list, 0)
- return X_embed_list, y_list
- def get_embed_domain(domain_dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, c_model) :
- fe_model_methyl.eval()
- fe_model_gene.eval()
- fe_model_multiomics.eval()
- c_model.eval()
- X_embed_list = []
- X_embed_methyl_list = []
- X_embed_gene_list = []
- domain_list = []
- pred_subtype_list = []
- label_list = [] # Can be used only for source dataset
- with torch.no_grad() :
- for batch, (X, X_gene, y, z) in enumerate(domain_dataloader):
- X, X_gene, y, z = X.to(device), X_gene.to(device), y.to(device), z.to(device)
- X = X.float()
- X_gene = X_gene.float()
- #
- embed_methyl = fe_model_methyl(X)
- embed_gene = fe_model_gene(X_gene)
- embed_concated = torch.cat((embed_methyl, embed_gene), 1)
- X_embed = fe_model_multiomics(embed_concated)
- #
- #X_embed = fe_model(X)
- pred = c_model(X_embed)
- pred_subtype_list.append(pred.argmax(1))
- X_embed_list.append(X_embed)
- domain_list.append(y)
- label_list.append(z)
- X_embed_methyl_list.append(embed_methyl)
- X_embed_gene_list.append(embed_gene)
- X_embed_list = torch.cat(X_embed_list, 0)
- X_embed_methyl_list = torch.cat(X_embed_methyl_list, 0)
- X_embed_gene_list = torch.cat(X_embed_gene_list, 0)
- pred_subtype_list = torch.cat(pred_subtype_list, 0)
- domain_list = torch.cat(domain_list, 0)
- label_list = torch.cat(label_list, 0)
- return X_embed_list, domain_list, pred_subtype_list, label_list, X_embed_methyl_list, X_embed_gene_list
- pt_epochs = 20#500
- ad_train_epochs = 20#500
- ssl_train_epochs = 20#500
- ft_epochs = 20#800
- # 1. Pre-training
- for t in range(pt_epochs):
- pretrain_classifier(t, train_dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, subtype_pred_model, c_loss, fe_methyl_optimizer, fe_gene_optimizer, fe_multiomics_optimizer, c_optimizer)
- # 2-1. Adversarial training (Single-omics)
- for t in range(ad_train_epochs):
- adversarial_train_disc_single_omics(t, domain_dataloader, fe_model_methyl, fe_model_gene, domain_disc_methyl_model, domain_disc_gene_model, domain_loss, fe_methyl_optimizer, fe_gene_optimizer, d_methyl_optimizer, d_gene_optimizer)
- adversarial_train_fe_single_omics(t, domain_dataloader, fe_model_methyl, fe_model_gene, domain_disc_methyl_model, domain_disc_gene_model, domain_loss, fe_methyl_optimizer, fe_gene_optimizer, d_methyl_optimizer, d_gene_optimizer)
- # 2-2. Adversarial training (Multiomics)
- for t in range(ad_train_epochs):
- adversarial_train_disc_multiomics(t, domain_dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, domain_disc_multiomics_model, domain_loss, fe_multiomics_optimizer, d_multiomics_optimizer)
- adversarial_train_fe_multiomics(t, domain_dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, domain_disc_multiomics_model, domain_loss, fe_multiomics_optimizer, d_multiomics_optimizer)
- # 3. SSL training
- for t in range(ssl_train_epochs) :
- target_pseudo_label = ssl_train_classifier(t, train_dataloader, target_dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, subtype_pred_model, c_loss, fe_methyl_optimizer, fe_gene_optimizer, fe_multiomics_optimizer, c_optimizer)
- target_dataset = SourceDataset(target_x, target_x_gene, target_pseudo_label)
- target_dataloader = DataLoader(target_dataset, batch_size = target_batch_size)
- # 4. Fine-tuning
- for t in range(ft_epochs) :
- # SSL
- target_pseudo_label = ssl_train_classifier(t, train_dataloader, target_dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, subtype_pred_model, c_loss, fe_methyl_optimizer, fe_gene_optimizer, fe_multiomics_optimizer, c_optimizer)
- target_dataset = SourceDataset(target_x, target_x_gene, target_pseudo_label)
- target_dataloader = DataLoader(target_dataset, batch_size = target_batch_size)
- # CA
- target_pseudo_label = target_pseudo_label.to("cpu")
- domain_z = torch.cat((y_train, target_pseudo_label), 0)
- domain_dataset = DomainDataset(domain_x, domain_x_gene, domain_y, domain_z)
- domain_dataloader = DataLoader(domain_dataset, batch_size = batch_size, shuffle = True)
- class_alignment_train(t, domain_dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, fe_methyl_optimizer, fe_gene_optimizer, fe_multiomics_optimizer, c_optimizer)
- domain_dataset = DomainDataset(domain_x, domain_x_gene, domain_y, domain_z)
- domain_dataloader = DataLoader(domain_dataset, batch_size = batch_size, shuffle = False)
- data_X_embed, domain_label, pred_subtype, label_subtype, methyl_embed, gene_embed = get_embed_domain(domain_dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, subtype_pred_model)
- data_X_embed = data_X_embed.detach().cpu().numpy()
- domain_label = domain_label.detach().cpu().numpy()
- pred_subtype = pred_subtype.detach().cpu().numpy()
- label_subtype = label_subtype.detach().cpu().numpy()
- #methyl_embed = methyl_embed.detach().cpu().numpy()
- #gene_embed = gene_embed.detach().cpu().numpy()
- data_X_embed = pd.DataFrame(data_X_embed)
- data_X_embed['Batch'] = domain_label
- data_X_embed['Pred_subtype'] = pred_subtype
- data_X_embed['Label_subtype'] = label_subtype
- data_X_embed.index = sample_id_list
- domain_info = pd.read_csv(os.path.join(sourceDataDir, "batch_category_info.csv"), index_col = 1)
- subtype_info = pd.read_csv(os.path.join(sourceDataDir, "subtype_category_info.csv"), index_col = 1)
- domain_info = domain_info.to_dict()
- domain_info['batch'][0] = 'Source'
- subtype_info = subtype_info.to_dict()
- data_X_embed['Pred_subtype'] = data_X_embed['Pred_subtype'].replace(subtype_info['subtype'])
- data_X_embed['Label_subtype'] = data_X_embed['Label_subtype'].replace(subtype_info['subtype'])
- data_X_embed['Batch'] = data_X_embed['Batch'].replace(domain_info['batch'])
- data_X_embed.to_csv(os.path.join(result_dir, "batch_corrected_features.csv"), mode = "w", index = True)
- target_pred = data_X_embed[['Batch','Pred_subtype']]
- target_pred = target_pred[target_pred['Batch'] != 'Source']
- target_pred.to_csv(os.path.join(result_dir, "results_target_prediction.csv"), mode = "w", index = True)
moDAmix.py at commit bd7e3c4, under MIT · at the source
Overview
- Department of Computer Science, Virginia Tech,Blacksburg, 24061 USA
- Division of Computer Science, Sookmyung Women’s University,Seoul, 04310 South Korea
Abstract
Multi-omics studies enable a comprehensive understanding of biological systems by integrating complementary molecular layers such as gene expression, DNA methylation, and chromatin accessibility. However, the generation of multi-omics data remains costly and labor-intensive, leading researchers to combine publicly available datasets collected from different cohorts, laboratories, and platforms. Integrating such heterogeneous datasets introduces substantial batch effects and technical variability that can obscure true biological structure. While numerous batch correction methods exist for single-omics data, systematic approaches for multi-omics batch effect correction remain limited. Correcting each omics layer independently risks disrupting cross-omics concordance and fails to ensure that samples are aligned within a unified multi-modal space, underscoring the need for coordinated, modality-aware harmonization that preserves shared molecular structure while removing technical variation across studies. To address this gap, we developed MoDAmix, a unified framework that leverages domain adaptation to remove technical variation while preserving shared molecular structure across omics layers. In particular, MoDAmix aligns feature distributions across batches and modalities through adversarial learning, enforcing consistency both within and between omics types to achieve coherent cross-omics integration. MoDAmix proceeds through four stages: (1) pre-training to learn initial feature representations, (2) adversarial adaptation to reduce batch effects within each omics type, (3) multi-omics adversarial alignment to harmonize modalities in a shared latent space, and (4) semi-supervised class alignment to refine subtype separability through pseudo-labeling and centroid consistency. Evaluations on both single-cell and bulk datasets–including mouse brain (gene expression and chromatin accessibility) and cancer cohorts (gene expression and DNA methylation)–demonstrate
Reproduced under the paper's license (CC BY), from the paper cited above.
Repository
Its files are read in the Code ↔ Paper reader above, with 2 matches between paragraphs and lines of code.
cbi-bioinfo/MoDAmix
bd7e3c4502c4c66813296c10ec7983e9104ac03f, 3 February 2026Availability: 1 check, the latest on 30 September 2026: the link answers
- 30 September 2026: the link answers
6 files
- __init__.py, Python, 1 line
- moDAmix.py, Python, 681 lines, 2 matches
- run_MoDAmix.sh, Shell, 12 lines
- setup.py, Python, 22 lines
- LICENSE, License, 21 lines
- README.md, Text, 69 lines
The paper's code and data availability statement is in the Data section.
Tracing map
Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.
What the map holds:
- 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
- 4 scripts, each with its path and the digest of its content;
- 2 matches between paragraphs of the paper and lines of the code (method lexical-v1);
- neither the text of the paper nor the code itself.
Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.
Data
Data links
- ncbi.nlm.nih.gov/
geo , NCBI; found in “Data availability”
Data availability
TCGA-LAML, TARGET-AML, TCGA-LGG, and CPTAC-3 datasets are available from GDC Data Portal (https://
Reproduced under the paper's license (CC BY), from the paper cited above.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 30 September 2026: the first record
Recorded: type, language, journal, volume, issue, pages, dates, 2 authors, 2 keywords, 6 MeSH terms, 2 funders, 29 references.
Cite
This paper
Choi, J. M., & Chae, H. (2026). A unified framework for correcting batch effects and integrating multi-omics data. Scientific reports, 16(1), 12341. https://
BibTeX
@article{choi2026unified
author = {Choi, Joung Min and Chae, Heejoon},
title = {{A unified framework for correcting batch effects and integrating multi-omics data}},
journal = {Scientific reports},
year = {2026},
month = mar,
volume = {16},
number = {1},
pages = {12341},
publisher = {Nature Publishing Group},
issn = {2045-2322},
doi = {10.1038/
url = {https://
pmid = {41786846},
pmcid = {PMC13079841}
}
RIS
TY - JOUR
AU - Choi, Joung Min
AU - Chae, Heejoon
TI - A unified framework for correcting batch effects and integrating multi-omics data
T2 - Scientific reports
J2 - Sci Rep
PY - 2026
DA - 2026/
VL - 16
IS - 1
SP - 12341
SN - 2045-2322
PB - Nature Publishing Group
DO - 10.1038/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1038/
"type": "article-journal",
"title": "A unified framework for correcting batch effects and integrating multi-omics data",
"container-title": "Scientific reports",
"author": [
{
"family": "Choi",
"given": "Joung Min"
},
{
"family": "Chae",
"given": "Heejoon"
}
],
"container-title-short":
"volume": "16",
"issue": "1",
"page": "12341",
"DOI": "10.1038/
"PMID": "41786846",
"PMCID": "PMC13079841",
"ISSN": "2045-2322",
"publisher": "Nature Publishing Group",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
3,
5
]
]
}
}
The tracing map gets a citation of its own once an author has validated it and it has a DOI.
Similar papers
The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.
- [1] doi:10.1038/s44320-026-00208-7 [code]
- Interpretable deep generative ensemble learning for single-cell omics with Hydra.Journal: Molecular systems biologyIn common: PyTorch, pandas, NumPy, 5 references
- [2] doi:10.1093/bioinformatics/btag652 [code]
- mmVelo: a deep generative model for estimating cell state-dependent dynamics across multiple modalities.Journal: Bioinformatics (Oxford, England)In common: PyTorch, pandas, NumPy, genetics / omics, 4 references
- [3] doi:10.1038/s41467-026-68596-w [code]
- Spatial cartography of human thymus enables the geopositioning of lineage transcription factors in rare mimetic thymic epithelial cells.Journal: Nature communicationsIn common: PyTorch, pandas, NumPy, genetics / omics, 2 references
- [4] doi:10.1371/journal.pone.0351405 [code]
- Unimodal vs. multimodal deep learning for non-invasive MGMT promoter methylation prediction in glioblastoma: A systematic evaluation on the BraTS 2021 dataset.Journal: PloS oneIn common: PyTorch, pandas, NumPy, methods / tools, genetics / omics, other condition, 1 reference
- [5] doi:10.1038/s41592-026-03057-2 [code]
- CREsted: modeling genomic and synthetic cell-type-specific enhancers across tissues and species.Journal: Nature methodsIn common: PyTorch, pandas, NumPy, methods / tools, genetics / omics, 2 references
- [6] doi:10.1002/advs.77003 [code]
- SemanticST: A Scalable Multi-Contextual Graph Learning Framework for Uncovering Spatial Niches and Robust Multi-Sample Integration in Spatial Transcriptomics.Journal: Advanced science (Weinheim, Baden-Wurttemberg, Germany)In common: PyTorch, pandas, NumPy, methods / tools, genetics / omics, 2 references
- [7] doi:10.1093/bioinformatics/btag253 [code]
- PEARL: integrative multi-omics classification and omics feature discovery via deep graph learning.Journal: Bioinformatics (Oxford, England)In common: PyTorch, pandas, NumPy, methods / tools, genetics / omics, 1 reference
- [8] doi:10.1016/j.xcrm.2026.102655 [code]
- An integrative multi-omics approach identifies microbiome alterations linked to pathological and behavioral features in autism spectrum disorder.Journal: Cell reports. MedicineIn common: pandas, NumPy, genetics / omics, 2 references
- [9] doi:10.3389/fneur.2026.1822479 [code]
- Circulating neuron-derived cfDNA for blood-based detection of Alzheimer's and other neurodegenerative conditions.Journal: Frontiers in neurologyIn common: PyTorch, pandas, NumPy, genetics / omics, other condition, 1 reference
- [10] doi:10.3389/fsysb.2026.1873899 [code]
- A systems microbiology framework for reproducible multi-dataset omics integration with application to long COVID.Journal: Frontiers in systems biologyIn common: PyTorch, pandas, NumPy, methods / tools, genetics / omics, 1 reference
Contribute
The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.
Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 4 scripts, and 2 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:4d553dffb8baeb7a…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
