OSCR

A unified framework for correcting batch effects and integrating multi-omics data.

Code ↔ Paper

2 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 2 matches
  1. [1] § Methods › Hyperparameter setting ↔ moDAmix.py, lines 225–246 · score 0.67 · domain discriminators, single omics feature, multi omics feature, Adam, softmax, PyTorch
  2. [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

  1. import torch
  2. from torch import nn
  3. from torch.utils.data import DataLoader
  4. from torch.utils.data import Dataset
  5. import pandas as pd
  6. import os
  7. import numpy as np
  8. import sys
  9. class SourceDataset(Dataset):
  10. def __init__(self, x_data, x_gene, y_data):
  11. self.x_data = x_data
  12. self.x_gene = x_gene
  13. self.y_data = y_data
  14. def __getitem__(self, index):
  15. return self.x_data[index], self.x_gene[index], self.y_data[index]
  16. def __len__(self):
  17. return self.x_data.shape[0]
  18. class UnlabelDataset(Dataset):
  19. def __init__(self, x_data, x_gene):
  20. self.x_data = x_data
  21. self.x_gene = x_gene
  22. def __getitem__(self, index):
  23. return self.x_data[index], self.x_gene[index]
  24. def __len__(self):
  25. return self.x_data.shape[0]
  26. class DomainDataset(Dataset) :
  27. def __init__(self, x_data, x_gene, y_data, z_data):
  28. self.x_data = x_data
  29. self.x_gene = x_gene
  30. self.y_data = y_data
  31. self.z_data = z_data
  32. def __getitem__(self, index):
  33. return self.x_data[index], self.x_gene[index], self.y_data[index], self.z_data[index]
  34. def __len__(self):
  35. return self.x_data.shape[0]
  36. device = (
  37. "cuda"
  38. if torch.cuda.is_available()
  39. else "cpu"
  40. )
  41. if device == "cuda" :
  42. os.environ["CUDA_VISIBLE_DEVICES"] = "0"
  43. print(f"Using {device} device")
  44. result_dir = "./results"
  45. os.makedirs(result_dir, exist_ok = True)
  46. data_dir = sys.argv[1]
  47. sourceDataDir = data_dir
  48. targetDataDir = data_dir
  49. x_filename = os.path.join(sourceDataDir, sys.argv[2])
  50. y_filename = os.path.join(sourceDataDir, sys.argv[4])
  51. target_filename = os.path.join(targetDataDir, sys.argv[5])
  52. x_gene_filename = os.path.join(sourceDataDir, sys.argv[3])
  53. target_gene_filename = os.path.join(targetDataDir, sys.argv[6])
  54. raw_x = pd.read_csv(x_filename, index_col = 0)
  55. raw_y = pd.read_csv(y_filename, index_col = 0)
  56. raw_target_x = pd.read_csv(target_filename, index_col = 0)
  57. raw_x_gene = pd.read_csv(x_gene_filename, index_col = 0)
  58. raw_target_x_gene = pd.read_csv(target_gene_filename, index_col = 0)
  59. sample_id_list = raw_x.index.tolist()
  60. sample_id_list.extend(raw_target_x.index.tolist())
  61. sample_id_list_gene = raw_x_gene.index.tolist()
  62. sample_id_list_gene.extend(raw_target_x_gene.index.tolist())
  63. raw_target_domain_y = raw_target_x['domain_idx'].tolist()
  64. raw_target_domain_y_gene = raw_target_x_gene['domain_idx'].tolist()
  65. raw_y_colname = raw_y.columns.tolist()[0]
  66. y_train = raw_y[raw_y_colname].tolist()
  67. num_subtype = len(set(y_train))
  68. y_train = np.array(y_train)
  69. del raw_target_x['domain_idx']
  70. del raw_target_x['Batch']
  71. del raw_target_x_gene['domain_idx']
  72. del raw_target_x_gene['Batch']
  73. raw_target_x = raw_target_x.values
  74. x_train = raw_x.values
  75. raw_target_x_gene = raw_target_x_gene.values
  76. x_train_gene = raw_x_gene.values
  77. domain_x = np.append(x_train, raw_target_x, axis = 0)
  78. domain_x_gene = np.append(x_train_gene, raw_target_x_gene, axis = 0)
  79. raw_source_domain_y = np.zeros(len(y_train), dtype = int) # TCGA label : 0
  80. domain_y = np.append(raw_source_domain_y, raw_target_domain_y)
  81. raw_source_domain_y_gene = np.zeros(len(raw_x_gene), dtype = int) # TCGA label : 0
  82. domain_y_gene = np.append(raw_source_domain_y_gene, raw_target_domain_y_gene)
  83. num_domain = len(set(domain_y))
  84. x_train = torch.from_numpy(x_train)
  85. y_train = torch.from_numpy(y_train)
  86. domain_x = torch.from_numpy(domain_x)
  87. domain_y = torch.from_numpy(domain_y)
  88. x_train_gene = torch.from_numpy(x_train_gene)
  89. domain_x_gene = torch.from_numpy(domain_x_gene)
  90. domain_y_gene = torch.from_numpy(domain_y_gene)
  91. target_x = torch.from_numpy(raw_target_x)
  92. target_x_gene = torch.from_numpy(raw_target_x_gene)
  93. target_init_y = torch.randint(low=0, high=num_subtype, size = (len(target_x),))
  94. #domain_z : domain_subtype
  95. domain_z = torch.cat((y_train, target_init_y), 0)
  96. num_feature = len(x_train[0])
  97. num_feature_gene = len(x_train_gene[0])
  98. num_train = len(x_train)
  99. num_test = len(raw_target_x)
  100. train_dataset = SourceDataset(x_train, x_train_gene, y_train)
  101. domain_dataset = DomainDataset(domain_x, domain_x_gene, domain_y, domain_z)
  102. target_dataset = SourceDataset(target_x, target_x_gene, target_init_y)
  103. batch_size = 128
  104. target_batch_size = 128
  105. test_target_batch_size = 64
  106. train_dataloader = DataLoader(train_dataset, batch_size = batch_size, shuffle = True)
  107. domain_dataloader = DataLoader(domain_dataset, batch_size = batch_size, shuffle = True)
  108. target_dataloader = DataLoader(target_dataset, batch_size = target_batch_size, shuffle = False)
  109. n_fe_embed1 = 1024
  110. n_fe_embed2 = 512
  111. n_mo_fe_embed1 = 512
  112. n_mo_fe_embed2 = 256
  113. n_c_h1 = 128
  114. n_c_h2 = 64
  115. n_d_h1 = 256
  116. n_d_h2 = 64
  117. class SingleOmicsFeatureExtractor(nn.Module) :
  118. def __init__(self, n_input) :
  119. super().__init__()
  120. self.feature_layer = nn.Sequential(
  121. nn.Linear(n_input, n_fe_embed1),
  122. nn.LeakyReLU(),
  123. nn.Linear(n_fe_embed1, n_fe_embed2),
  124. nn.LeakyReLU()
  125. )
  126. def forward(self, x) :
  127. embedding = self.feature_layer(x)
  128. return embedding
  129. class MultiOmicsFeatureExtractor(nn.Module) :
  130. def __init__(self) :
  131. super().__init__()
  132. self.feature_layer = nn.Sequential(
  133. nn.Linear(n_fe_embed2*2, n_mo_fe_embed1),
  134. nn.LeakyReLU(),
  135. nn.Linear(n_mo_fe_embed1, n_mo_fe_embed2),
  136. nn.LeakyReLU()
  137. )
  138. def forward(self, x) :
  139. embedding = self.feature_layer(x)
  140. return embedding
  141. class DomainDiscriminator(nn.Module) :
  142. def __init__(self, n_fe_h2) :
  143. super().__init__()
  144. self.disc_layer = nn.Sequential(
  145. nn.Linear(n_fe_h2, n_d_h1),
  146. nn.LeakyReLU(),
  147. nn.Linear(n_d_h1, n_d_h2),
  148. nn.LeakyReLU(),
  149. nn.Linear(n_d_h2, num_domain)
  150. )
  151. def forward(self, x) :
  152. domain_logits = self.disc_layer(x)
  153. return domain_logits
  154. class SubtypeClassifier(nn.Module):
  155. def __init__(self):
  156. super().__init__()
  157. #self.flatten = nn.Flatten()
  158. self.linear_relu_stack = nn.Sequential(
  159. nn.Linear(n_mo_fe_embed2, n_c_h1),
  160. nn.LeakyReLU(),
  161. nn.Linear(n_c_h1, n_c_h2),
  162. nn.LeakyReLU(),
  163. nn.Linear(n_c_h2, num_subtype)
  164. )
  165. def forward(self, x):
  166. logits = self.linear_relu_stack(x)
  167. return logits
  168. fe_model_methyl = SingleOmicsFeatureExtractor(num_feature).to(device)
  169. fe_model_gene = SingleOmicsFeatureExtractor(num_feature_gene).to(device)
  170. fe_model_multiomics = MultiOmicsFeatureExtractor().to(device)
  171. domain_disc_methyl_model = DomainDiscriminator(n_fe_embed2).to(device)
  172. domain_disc_gene_model = DomainDiscriminator(n_fe_embed2).to(device)
  173. domain_disc_multiomics_model = DomainDiscriminator(n_mo_fe_embed2).to(device)
  174. subtype_pred_model = SubtypeClassifier().to(device)
  175. c_loss = nn.CrossEntropyLoss() # Already have softmax
  176. domain_loss = nn.CrossEntropyLoss() # Already have softmax
  177. fe_methyl_optimizer = torch.optim.Adam(fe_model_methyl.parameters(), lr=1e-4)
  178. fe_gene_optimizer = torch.optim.Adam(fe_model_gene.parameters(), lr=1e-4)
  179. fe_multiomics_optimizer = torch.optim.Adam(fe_model_multiomics.parameters(), lr=1e-4)
  180. c_optimizer = torch.optim.Adam(subtype_pred_model.parameters(), lr=1e-5)
  181. d_methyl_optimizer = torch.optim.Adam(domain_disc_methyl_model.parameters(), lr=1e-6)
  182. d_gene_optimizer = torch.optim.Adam(domain_disc_gene_model.parameters(), lr=1e-6)
  183. d_multiomics_optimizer = torch.optim.Adam(domain_disc_multiomics_model.parameters(), lr=1e-6)
  184. 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):
  185. size = len(dataloader.dataset)
  186. correct = 0
  187. for batch, (X, X_gene, y) in enumerate(dataloader):
  188. X, X_gene, y = X.to(device), X_gene.to(device), y.to(device)
  189. X = X.float()
  190. X_gene = X_gene.float()
  191. embed_methyl = fe_model_methyl(X)
  192. embed_gene = fe_model_gene(X_gene)
  193. embed_concated = torch.cat((embed_methyl, embed_gene), 1)
  194. embed_multiomics = fe_model_multiomics(embed_concated)
  195. pred = c_model(embed_multiomics)
  196. loss = c_loss(pred, y)
  197. fe_methyl_optimizer.zero_grad()
  198. fe_gene_optimizer.zero_grad()
  199. fe_multiomics_optimizer.zero_grad()
  200. c_optimizer.zero_grad()
  201. loss.backward()
  202. fe_methyl_optimizer.step()
  203. fe_gene_optimizer.step()
  204. fe_multiomics_optimizer.step()
  205. c_optimizer.step()
  206. correct += (pred.argmax(1) == y).type(torch.float).sum().item()
  207. loss = loss.item()
  208. correct /= size
  209. if epoch % 10 == 0 :
  210. print(f"[PT Epoch {epoch+1}] \tTraining loss: {loss:>5f}, Training Accuracy: {(100*correct):>0.2f}%")
  211. 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) :
  212. size = len(dataloader.dataset)
  213. correct_methyl = 0
  214. correct_gene = 0
  215. for batch, (X, X_gene, y, z_subtype) in enumerate(dataloader):
  216. X, X_gene, y = X.to(device), X_gene.to(device), y.to(device)
  217. X = X.float()
  218. X_gene = X_gene.float()
  219. embed_methyl = fe_model_methyl(X)
  220. pred = d_model_methyl(embed_methyl)
  221. d_loss = domain_loss(pred, y)
  222. # Backpropagation for methyl
  223. fe_methyl_optimizer.zero_grad()
  224. d_methyl_optimizer.zero_grad()
  225. d_loss.backward()
  226. d_methyl_optimizer.step()
  227. correct_methyl += (pred.argmax(1) == y).type(torch.float).sum().item()
  228. #
  229. embed_gene = fe_model_gene(X_gene)
  230. pred_gene = d_model_gene(embed_gene)
  231. d_gene_loss = domain_loss(pred_gene, y)
  232. fe_gene_optimizer.zero_grad()
  233. d_gene_optimizer.zero_grad()
  234. d_gene_loss.backward()
  235. d_gene_optimizer.step()
  236. correct_gene += (pred_gene.argmax(1) == y).type(torch.float).sum().item()
  237. d_loss = d_loss.item()
  238. d_gene_loss = d_gene_loss.item()
  239. correct_methyl /= size
  240. correct_gene /= size
  241. if t % 10 == 0 :
  242. 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 = ", ")
  243. 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) :
  244. size = len(dataloader.dataset)
  245. correct = 0
  246. for batch, (X, X_gene, y, z_subtype) in enumerate(dataloader):
  247. X, X_gene, y = X.to(device), X_gene.to(device), y.to(device)
  248. X = X.float()
  249. X_gene = X_gene.float()
  250. #
  251. embed_methyl = fe_model_methyl(X)
  252. embed_gene = fe_model_gene(X_gene)
  253. embed_concated = torch.cat((embed_methyl, embed_gene), 1)
  254. embed_multiomics = fe_model_multiomics(embed_concated)
  255. #
  256. pred = d_model_multiomics(embed_multiomics)
  257. d_loss = domain_loss(pred, y)
  258. # Backpropagation
  259. fe_multiomics_optimizer.zero_grad()
  260. d_multiomics_optimizer.zero_grad()
  261. d_loss.backward()
  262. d_multiomics_optimizer.step()
  263. correct += (pred.argmax(1) == y).type(torch.float).sum().item()
  264. d_loss = d_loss.item()
  265. correct /= size
  266. if t % 10 == 0 :
  267. print(f"[AT-M Epoch {epoch+1}] Disc loss: {d_loss:>5f}, Training Accuracy: {(100*correct):>0.2f}%", end = ", ")
  268. 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) :
  269. size = len(dataloader.dataset)
  270. for batch, (X, X_gene, y, z_subtype) in enumerate(dataloader):
  271. X, X_gene, y = X.to(device), X_gene.to(device), y.to(device)
  272. X = X.float()
  273. X_gene = X_gene.float()
  274. embed_methyl = fe_model_methyl(X)
  275. pred = d_model_methyl(embed_methyl)
  276. fake_y = torch.randint(low=0, high=num_domain, size = (len(y),))
  277. fake_y = fake_y.to(device)
  278. g_loss = domain_loss(pred, fake_y)
  279. # Backpropagation
  280. fe_methyl_optimizer.zero_grad()
  281. d_methyl_optimizer.zero_grad()
  282. g_loss.backward()
  283. fe_methyl_optimizer.step()
  284. # Gene
  285. embed_gene = fe_model_gene(X_gene)
  286. pred_gene = d_model_gene(embed_gene)
  287. fake_y_gene = torch.randint(low=0, high=num_domain, size = (len(y),))
  288. fake_y_gene = fake_y_gene.to(device)
  289. g_gene_loss = domain_loss(pred_gene, fake_y)
  290. fe_gene_optimizer.zero_grad()
  291. d_gene_optimizer.zero_grad()
  292. g_gene_loss.backward()
  293. fe_gene_optimizer.step()
  294. g_loss = g_loss.item()
  295. g_gene_loss = g_gene_loss.item()
  296. if epoch % 10 == 0:
  297. print(f"Gen methyl loss: {g_loss:>5f}, Gene gene loss: {g_gene_loss:>5f}")
  298. 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) :
  299. size = len(dataloader.dataset)
  300. for batch, (X, X_gene, y, z_subtype) in enumerate(dataloader):
  301. X, X_gene, y = X.to(device), X_gene.to(device), y.to(device)
  302. X = X.float()
  303. X_gene = X_gene.float()
  304. #
  305. embed_methyl = fe_model_methyl(X)
  306. embed_gene = fe_model_gene(X_gene)
  307. embed_concated = torch.cat((embed_methyl, embed_gene), 1)
  308. embed_multiomics = fe_model_multiomics(embed_concated)
  309. #
  310. pred = d_model_multiomics(embed_multiomics)
  311. fake_y = torch.randint(low=0, high=num_domain, size = (len(y),))
  312. fake_y = fake_y.to(device)
  313. g_loss = domain_loss(pred, fake_y)
  314. # Backpropagation
  315. fe_multiomics_optimizer.zero_grad()
  316. d_multiomics_optimizer.zero_grad()
  317. g_loss.backward()
  318. fe_multiomics_optimizer.step()
  319. g_loss = g_loss.item()
  320. if epoch % 10 == 0:
  321. print(f"Gen multi loss: {g_loss:>5f}")
  322. 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) :
  323. for batch, (X, X_gene, y_domain, z_subtype) in enumerate(domain_dataloader):
  324. X, X_gene, y_domain, z_subtype = X.to(device), X_gene.to(device), y_domain.to(device), z_subtype.to(device)
  325. X = X.float()
  326. X_gene = X_gene.float()
  327. batch_subtype_list = z_subtype.unique()
  328. #
  329. embed_methyl = fe_model_methyl(X)
  330. embed_gene = fe_model_gene(X_gene)
  331. embed_concated = torch.cat((embed_methyl, embed_gene), 1)
  332. X_embed = fe_model_multiomics(embed_concated)
  333. #
  334. align_loss = torch.zeros((1) ,dtype = torch.float64)
  335. align_loss = align_loss.to(device)
  336. #
  337. for subtype in batch_subtype_list :
  338. sample_idx_list = (z_subtype == subtype).nonzero(as_tuple = True)[0]
  339. if len(sample_idx_list) < 1 :
  340. continue
  341. #else :
  342. tmp_x = X_embed[sample_idx_list]
  343. tmp_y = y_domain[sample_idx_list]
  344. tmp_z = z_subtype[sample_idx_list]
  345. batch_domain_list = tmp_y.unique()
  346. domain_centroid_stack = []
  347. for domain in batch_domain_list :
  348. domain_idx_list = (tmp_y == domain).nonzero(as_tuple = True)[0]
  349. if len(domain_idx_list) != 1 :
  350. tmp_x_domain = tmp_x[domain_idx_list]
  351. tmp_centroid = torch.div(torch.sum(tmp_x_domain, dim = 0), len(domain_idx_list))
  352. domain_centroid_stack.append(tmp_centroid)
  353. if len(domain_centroid_stack) == 0 :
  354. continue
  355. else :
  356. domain_centroid_stack = torch.stack(domain_centroid_stack)
  357. subtype_centroid = torch.mean(domain_centroid_stack, dim = 0)
  358. # Duplicate the subtype centroid to get dist with each domain_centroid
  359. subtype_centroid_stack = []
  360. for i in range(len(domain_centroid_stack)) :
  361. subtype_centroid_stack.append(subtype_centroid)
  362. subtype_centroid_stack = torch.stack(subtype_centroid_stack)
  363. pdist_stack = nn.L1Loss()(subtype_centroid_stack, domain_centroid_stack)
  364. align_loss += torch.mean(pdist_stack, dim = 0)
  365. if align_loss == 0.0 :
  366. continue
  367. align_loss = align_loss / len(batch_subtype_list)
  368. fe_methyl_optimizer.zero_grad()
  369. fe_gene_optimizer.zero_grad()
  370. fe_multiomics_optimizer.zero_grad()
  371. c_optimizer.zero_grad()
  372. align_loss.backward()
  373. fe_methyl_optimizer.step()
  374. fe_gene_optimizer.step()
  375. fe_multiomics_optimizer.step()
  376. c_optimizer.step()
  377. align_loss = align_loss.item()
  378. if epoch % 10 == 0 :
  379. print(f"[CA Epoch {epoch+1}] align loss: {align_loss:>5f}\n")
  380. 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) :
  381. source_size = len(source_dataloader.dataset)
  382. target_size = len(target_dataloader.dataset)
  383. #
  384. # 1. Obtain the pseudo-label for target dataset
  385. #
  386. target_pseudo_label = torch.empty((0), dtype = torch.int64)
  387. target_pseudo_label = target_pseudo_label.to(device)
  388. #
  389. for batch, (target_X, target_X_gene, target_y) in enumerate(target_dataloader):
  390. target_X, target_X_gene, target_y = target_X.to(device), target_X_gene.to(device), target_y.to(device)
  391. target_X = target_X.float()
  392. target_X_gene = target_X_gene.float()
  393. #
  394. embed_methyl = fe_model_methyl(target_X)
  395. embed_gene = fe_model_gene(target_X_gene)
  396. embed_concated = torch.cat((embed_methyl, embed_gene), 1)
  397. extracted_feature = fe_model_multiomics(embed_concated)
  398. #
  399. #extracted_feature = fe_model(target_X)
  400. batch_target_pred = c_model(extracted_feature)
  401. batch_pseudo_label = batch_target_pred.argmax(1)
  402. target_pseudo_label = torch.cat((target_pseudo_label, batch_pseudo_label), 0)
  403. if batch == 0 :
  404. target_loss = c_loss(batch_target_pred, target_y)
  405. else :
  406. target_loss = target_loss + c_loss(batch_target_pred, target_y)
  407. target_loss = target_loss / (batch + 1)
  408. #
  409. # Define alpha value
  410. alpha_f = 0.01
  411. t1 = 100
  412. t2 = 200
  413. if epoch < t1 :
  414. alpha = 0
  415. elif epoch < t2 :
  416. alpha = (epoch - t1) / (t2 - t1) * alpha_f
  417. else :
  418. alpha = alpha_f
  419. #
  420. # 2. Calculate the loss for the source dataset
  421. #
  422. correct = 0
  423. for batch, (source_X, source_X_gene, source_y) in enumerate(source_dataloader):
  424. source_X, source_X_gene, source_y = source_X.to(device), source_X_gene.to(device), source_y.to(device)
  425. source_X = source_X.float()
  426. source_X_gene = source_X_gene.float()
  427. #
  428. source_embed_methyl = fe_model_methyl(source_X)
  429. source_embed_gene = fe_model_gene(source_X_gene)
  430. source_embed_concated = torch.cat((source_embed_methyl, source_embed_gene), 1)
  431. source_extracted_feature = fe_model_multiomics(source_embed_concated)
  432. #
  433. #source_extracted_feature = fe_model(source_X)
  434. source_pred = c_model(source_extracted_feature)
  435. source_loss = c_loss(source_pred, source_y)
  436. ssl_loss = source_loss + alpha * target_loss
  437. # Backpropogation
  438. target_loss.detach_()
  439. fe_methyl_optimizer.zero_grad()
  440. fe_gene_optimizer.zero_grad()
  441. fe_multiomics_optimizer.zero_grad()
  442. c_optimizer.zero_grad()
  443. ssl_loss.backward() #retain_graph=True
  444. fe_methyl_optimizer.step()
  445. fe_gene_optimizer.step()
  446. fe_multiomics_optimizer.step()
  447. c_optimizer.step()
  448. correct += (source_pred.argmax(1) == source_y).type(torch.float).sum().item()
  449. ssl_loss = ssl_loss.item()
  450. source_loss = source_loss.item()
  451. target_loss = target_loss.item()
  452. correct /= source_size
  453. if epoch % 10 == 0 :
  454. 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")
  455. return target_pseudo_label
  456. def get_embed(dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, c_model) :
  457. fe_model_methyl.eval()
  458. fe_model_gene.eval()
  459. fe_model_multiomics.eval()
  460. c_model.eval()
  461. X_embed_list = []
  462. y_list = []
  463. with torch.no_grad() :
  464. for batch, (X, X_gene, y) in enumerate(dataloader):
  465. X, X_gene, y = X.to(device), X_gene.to(device), y.to(device)
  466. X = X.float()
  467. #
  468. embed_methyl = fe_model_methyl(X)
  469. embed_gene = fe_model_gene(X_gene)
  470. embed_concated = torch.cat((embed_methyl, embed_gene), 1)
  471. X_embed = fe_model_multiomics(embed_concated)
  472. #
  473. X_embed_list.append(X_embed)
  474. y_list.append(y)
  475. X_embed_list = torch.cat(X_embed_list, 0)
  476. y_list = torch.cat(y_list, 0)
  477. return X_embed_list, y_list
  478. def get_embed_domain(domain_dataloader, fe_model_methyl, fe_model_gene, fe_model_multiomics, c_model) :
  479. fe_model_methyl.eval()
  480. fe_model_gene.eval()
  481. fe_model_multiomics.eval()
  482. c_model.eval()
  483. X_embed_list = []
  484. X_embed_methyl_list = []
  485. X_embed_gene_list = []
  486. domain_list = []
  487. pred_subtype_list = []
  488. label_list = [] # Can be used only for source dataset
  489. with torch.no_grad() :
  490. for batch, (X, X_gene, y, z) in enumerate(domain_dataloader):
  491. X, X_gene, y, z = X.to(device), X_gene.to(device), y.to(device), z.to(device)
  492. X = X.float()
  493. X_gene = X_gene.float()
  494. #
  495. embed_methyl = fe_model_methyl(X)
  496. embed_gene = fe_model_gene(X_gene)
  497. embed_concated = torch.cat((embed_methyl, embed_gene), 1)
  498. X_embed = fe_model_multiomics(embed_concated)
  499. #
  500. #X_embed = fe_model(X)
  501. pred = c_model(X_embed)
  502. pred_subtype_list.append(pred.argmax(1))
  503. X_embed_list.append(X_embed)
  504. domain_list.append(y)
  505. label_list.append(z)
  506. X_embed_methyl_list.append(embed_methyl)
  507. X_embed_gene_list.append(embed_gene)
  508. X_embed_list = torch.cat(X_embed_list, 0)
  509. X_embed_methyl_list = torch.cat(X_embed_methyl_list, 0)
  510. X_embed_gene_list = torch.cat(X_embed_gene_list, 0)
  511. pred_subtype_list = torch.cat(pred_subtype_list, 0)
  512. domain_list = torch.cat(domain_list, 0)
  513. label_list = torch.cat(label_list, 0)
  514. return X_embed_list, domain_list, pred_subtype_list, label_list, X_embed_methyl_list, X_embed_gene_list
  515. pt_epochs = 20#500
  516. ad_train_epochs = 20#500
  517. ssl_train_epochs = 20#500
  518. ft_epochs = 20#800
  519. # 1. Pre-training
  520. for t in range(pt_epochs):
  521. 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)
  522. # 2-1. Adversarial training (Single-omics)
  523. for t in range(ad_train_epochs):
  524. 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)
  525. 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)
  526. # 2-2. Adversarial training (Multiomics)
  527. for t in range(ad_train_epochs):
  528. 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)
  529. 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)
  530. # 3. SSL training
  531. for t in range(ssl_train_epochs) :
  532. 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)
  533. target_dataset = SourceDataset(target_x, target_x_gene, target_pseudo_label)
  534. target_dataloader = DataLoader(target_dataset, batch_size = target_batch_size)
  535. # 4. Fine-tuning
  536. for t in range(ft_epochs) :
  537. # SSL
  538. 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)
  539. target_dataset = SourceDataset(target_x, target_x_gene, target_pseudo_label)
  540. target_dataloader = DataLoader(target_dataset, batch_size = target_batch_size)
  541. # CA
  542. target_pseudo_label = target_pseudo_label.to("cpu")
  543. domain_z = torch.cat((y_train, target_pseudo_label), 0)
  544. domain_dataset = DomainDataset(domain_x, domain_x_gene, domain_y, domain_z)
  545. domain_dataloader = DataLoader(domain_dataset, batch_size = batch_size, shuffle = True)
  546. 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)
  547. domain_dataset = DomainDataset(domain_x, domain_x_gene, domain_y, domain_z)
  548. domain_dataloader = DataLoader(domain_dataset, batch_size = batch_size, shuffle = False)
  549. 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)
  550. data_X_embed = data_X_embed.detach().cpu().numpy()
  551. domain_label = domain_label.detach().cpu().numpy()
  552. pred_subtype = pred_subtype.detach().cpu().numpy()
  553. label_subtype = label_subtype.detach().cpu().numpy()
  554. #methyl_embed = methyl_embed.detach().cpu().numpy()
  555. #gene_embed = gene_embed.detach().cpu().numpy()
  556. data_X_embed = pd.DataFrame(data_X_embed)
  557. data_X_embed['Batch'] = domain_label
  558. data_X_embed['Pred_subtype'] = pred_subtype
  559. data_X_embed['Label_subtype'] = label_subtype
  560. data_X_embed.index = sample_id_list
  561. domain_info = pd.read_csv(os.path.join(sourceDataDir, "batch_category_info.csv"), index_col = 1)
  562. subtype_info = pd.read_csv(os.path.join(sourceDataDir, "subtype_category_info.csv"), index_col = 1)
  563. domain_info = domain_info.to_dict()
  564. domain_info['batch'][0] = 'Source'
  565. subtype_info = subtype_info.to_dict()
  566. data_X_embed['Pred_subtype'] = data_X_embed['Pred_subtype'].replace(subtype_info['subtype'])
  567. data_X_embed['Label_subtype'] = data_X_embed['Label_subtype'].replace(subtype_info['subtype'])
  568. data_X_embed['Batch'] = data_X_embed['Batch'].replace(domain_info['batch'])
  569. data_X_embed.to_csv(os.path.join(result_dir, "batch_corrected_features.csv"), mode = "w", index = True)
  570. target_pred = data_X_embed[['Batch','Pred_subtype']]
  571. target_pred = target_pred[target_pred['Batch'] != 'Source']
  572. 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

Authors: Joung Min Choi1, Heejoon Chae2
  1. Department of Computer Science, Virginia Tech,Blacksburg, 24061 USA
  2. Division of Computer Science, Sookmyung Women’s University,Seoul, 04310 South Korea
Institutions: Virginia Tech (United States); Sookmyung Women's University (South Korea)
Journal: Scientific reports, volume 16, issue 1, article 12341
Dates: received 17 November 2025; accepted 25 February 2026; published online 5 March 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1038/s41598-026-42355-9 · PMID 41786846 · PMCID PMC13079841 · OpenAlex W7133904935
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: genetics / omics (modality), human (organism), other condition (population), methods / tools (subfield)
Methods: Smoothing, state filtering, decompositions, Machine learning
Keywords: Cancer, Computational biology and bioinformatics
MeSH: Computational Biology*, Genomics*, Multiomics*, Animals, DNA Methylation, Humans (* major topic)
Topic: Single-cell and spatial transcriptomics (Molecular Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Funding: Korea National Institute of Health (KNIH) research project (2024-ER-0801-01); Bio&Medical Technology Development Program of the National Research Foundation (NRF) funded by the Korean government (MSIT) (RS-2025-18732993)
Citations: cited by 4 papers (Europe PMC); 33 references in the paper

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)–demonstrated that MoDAmix effectively mitigates batch effects, improves clustering and classification performance, and preserves subtype structure across domains. Together, these results highlight MoDAmix as a robust framework for multi-omics batch effect correction and integration, enabling reliable cross-cohort analysis in systems biology and precision medicine. MoDAmix is publicly available at https://github.com/cbi-bioinfo/MoDAmix.

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

License: MIT
State: the link answers, verified on 30 September 2026
Evidence: files inventoried
Commit: bd7e3c4502c4c66813296c10ec7983e9104ac03f, 3 February 2026
Languages: Python (3), Shell (1)
Size: 25 files, 4 scripts
Software Heritage: not archived
Found in: “Data availability”
Holds: README, license file, environment (requirements.txt, setup.py)
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (2 files), pandas (2 files), PyTorch (2 files)
Availability: 1 check, the latest on 30 September 2026: the link answers
  • 30 September 2026: the link answers
6 files

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

Data availability

TCGA-LAML, TARGET-AML, TCGA-LGG, and CPTAC-3 datasets are available from GDC Data Portal (https://portal.gdc.cancer.gov/). Adult mouse brain dataset used in this study are available from GEO repository (https://www.ncbi.nlm.nih.gov/geo/) with the GEO accession of GSE130399 and GSE126074. Source codes of MoDAmix is publicly available at https://github.com/cbi-bioinfo/MoDAmix.

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://doi.org/10.1038/s41598-026-42355-9

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/s41598-026-42355-9},
url = {https://doi.org/10.1038/s41598-026-42355-9},
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/03/05
VL - 16
IS - 1
SP - 12341
SN - 2045-2322
PB - Nature Publishing Group
DO - 10.1038/s41598-026-42355-9
UR - https://doi.org/10.1038/s41598-026-42355-9
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41598-026-42355-9",
"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": "Sci Rep",
"volume": "16",
"issue": "1",
"page": "12341",
"DOI": "10.1038/s41598-026-42355-9",
"PMID": "41786846",
"PMCID": "PMC13079841",
"ISSN": "2045-2322",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41598-026-42355-9",
"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 biology
In 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 communications
In 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 one
In 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 methods
In 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. Medicine
In 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 neurology
In 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 biology
In 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.

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.