OSCR

Taming dimensionality in big neuroimaging data: Efficient orthonormal projective NMF via stochastic learning and data compression.

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] § Materials and Methods › Proposed methodology › Compressed sopNMF (csopNMF) ↔ Python/csopnmf.py, lines 86–196 · score 0.60 · QR decomposition, power iterations, target rank, mini batches, uniformly, compressed
  2. [2] § Materials and Methods › Experimental setup › Model validation on OASIS ↔ Python/initialize_nmf.py, lines 44–184 · score 0.51 · double singular, faster, NNDSVD, factorization, rank

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 · 1,469 lines · 73 KB · Apache-2.0 · 1 match

  1. ## Module Imports
  2. import sys #check python version, exit upon sanity check failures
  3. import os #path checks
  4. import fractions #for handling fractions in argparse
  5. import argparse #for taking in user input/parameter/switch
  6. import time #for measuring elapsed time
  7. import datetime #for default output path string formation
  8. import getpass #for getting username on compute jobs since os.getlogin() works in interactive jobs but not compute jobs
  9. import shutil #for renaming files
  10. import numpy as np #matrix multiplications and optimizations
  11. import pandas as pd #for reading csv
  12. import hdf5storage #for saving as hdf5 mat file (v7.3 on matlab); has much better compression than scipy savemat
  13. import utils #for loading and saving data
  14. import initialize_nmf #for W component initialization
  15. import opnmf_update_rule #for W component initialization
  16. import restart #for determining outdir
  17. ## Default Variables
  18. #default value for script name to be used for printing to console
  19. script_name=os.path.basename(__file__)
  20. script_dirname = os.path.dirname(__file__)
  21. username=str(getpass.getuser())
  22. #Version
  23. script_version = "20230911_160100"
  24. #sanity check: python version
  25. utils.check_python_version(major = 3, minor = 7) #ensure the user is using python3.7 or higher
  26. #FLAGS
  27. VERBOSE_FLAG = False
  28. DEBUG_FLAG = False
  29. WITH_REPLACEMENT_FLAG = False
  30. ## Default values to functions and argparser
  31. csv_path = os.path.join("/scratch/sungminha/git/NMF_Testing_Preprocessing/output_directory/dramms_1000_subjectid_only_those_without_negative_values_or_bad_ICV_changes_fullpaths_subjectid_age_sex_cdr_with_header.csv")
  32. feret_path = os.path.join("/scratch/sungminha/git/NMF_Testing/faces_1024_2409_min0_max1_double.mat")
  33. EPSILON = np.finfo(np.float32).eps
  34. ## cs-opNMF specific functions
  35. def left_compression_Q(X, k, k_ov, w):
  36. """
  37. X = data matrix
  38. k = target rank
  39. k_ov = oversampling, s.t. k + k_ov = l
  40. w = power iteration
  41. """
  42. m = X.shape[0]
  43. n = X.shape[1]
  44. omega = np.random.default_rng().normal(loc = 0, scale = 1, size = (n, l))
  45. B = np.matmul(X, omega)
  46. print("%s: left_compression_Q: B.shape = " % (script_name))
  47. print(B.shape, flush=True)
  48. for power_iteration in np.arange(w):
  49. B = np.matmul(X, np.matmul(X.T, B))
  50. Q, R = np.linalg.qr(B)
  51. return Q
  52. def right_compression_Q(X, k, k_ov, w):
  53. """
  54. X = data matrix
  55. k = target rank
  56. k_ov = oversampling, s.t. k + k_ov = l
  57. w = power iteration
  58. """
  59. m = X.shape[0]
  60. n = X.shape[1]
  61. omega = np.random.default_rng().normal(loc = 0, scale = 1, size = (l, m))
  62. B = np.matmul(omega, X)
  63. print("%s: right_compression_Q: B.shape = " % (script_name))
  64. print(B.shape, flush=True)
  65. for power_iteration in np.arange(w):
  66. B = np.matmul(np.matmul(B, X.T), X)
  67. Q, R = np.linalg.qr(B.T)
  68. return Q.T
  69. def get_Q_compression(X_hat, X_compression_size = 100, power_iteration = 4, l = 50, axis = 1,VERBOSE_FLAG = False ):
  70. """
  71. axis = 0: Get Q such that it compresses feature space of X: X[m,n] to X[l,n]
  72. axis = 1: Get Q such that it compresses subject space of X: X[m,n] to X[m,l]
  73. where we assume:
  74. l = np.amin([n, np.amax([compression_level, target_rank + 10])]) #r + r_ov
  75. X_compression_size: How many subjects of X to use for calculating Q
  76. e.g.) X[m,n], X_compression_size = n: then use full X
  77. e.g.) X[m,n], X_compression_Size = 0.1 * n: use X_hat[m,1/10*n] where we sample uniformly 1/10 of n subjects from X for calculating X
  78. recommended to use full X by providing X_compression_size = n if X matrix is not too large for computational resources
  79. l: output dimension after compression for the axis of interest
  80. axis = 0: X[m,n] -> X_compressed[l,n]
  81. axis = 1: X[m,n] -> X_compressed[m,l]
  82. axis = 0:
  83. Q = np.random.rand(m, l) shape of Q must be m by l
  84. omega = np.random.standard_normal(size = (n, l))
  85. axis = 1:
  86. Q = np.random.rand(l, n)
  87. omega = np.random.standard_normal(size = (l, m))
  88. """
  89. #sanity check: axis parameter must be 0 or 1 assuming 2D input data matrix
  90. if axis != 0 and axis != 1:
  91. print("%s: get_Q_compression: axis parameter must be 0 or 1, not %0.5E. Exiting." % (script_name, axis), flush=True)
  92. sys.exit(1)
  93. #sanity check: X must be 2D
  94. X_hat_shape = np.shape(X_hat)
  95. if np.shape(X_hat_shape)[0] != 2:
  96. print("%s: get_Q_compression: data matrix provided to calculate X_hat must be 2-D, but it is %i-D. Exiting." % (script_name, np.shape(X_hat_shape)[0]), flush=True)
  97. sys.exit(1)
  98. m = np.shape(X_hat)[0]
  99. n = np.shape(X_hat)[1]
  100. if VERBOSE_FLAG:
  101. print("%s: get_Q_compression: data matrix provided to calculate X_hat.shape = [%i, %i]." % (script_name, m, n), flush=True)
  102. #generate omega matrix accordingly to axis parameter
  103. if axis == 0:
  104. omega = np.random.default_rng().normal(loc = 0, scale = 1, size = (X_compression_size, l))
  105. elif axis == 1:
  106. omega = np.random.default_rng().normal(loc = 0, scale = 1, size = (l, m))
  107. else:
  108. print("%s: get_Q_compression: unknown value ( %0.5E ) for axis parameter. Exiting." % (script_name, axis),flush=True)
  109. sys.exit(1)
  110. if VERBOSE_FLAG:
  111. print("%s: omega.shape = " % (script_name), flush=True)
  112. print(omega.shape, flush=True)
  113. #check whether to calculate Q from mini-batch of X or full X, or exit if not a reasonable number of subjects to sample mini-batch is provided
  114. if VERBOSE_FLAG:
  115. print("%s: get_Q_compression: X_compression_size = %i" % (script_name, X_compression_size), flush=True)
  116. if n < X_compression_size:
  117. print("%s: get_Q_compression: X_compression_size ( %i ) must be less than or equal to number of columns of X_hat n ( %i ). Exiting." % (script_name, X_compression_size, n), flush=True)
  118. sys.exit(1)
  119. elif n == X_compression_size:
  120. X_compression_sample = X_hat
  121. else: #n > X_compression_size - need to sample mini-batch from X_hat
  122. idx_list = np.arange(start = 0, stop = n, step = 1) #list of indices to sample from
  123. X_compression_idx = np.random.choice(idx_list, size = X_compression_size, replace=False) #p=None default -> uniform random sampling
  124. X_compression_sample = X_hat[:,X_compression_idx]
  125. del X_compression_idx
  126. if VERBOSE_FLAG:
  127. print("%s: X_compression_sample.shape = " % (script_name), flush=True)
  128. print(X_compression_sample.shape, flush=True)
  129. if axis == 0:
  130. B = np.matmul(X_compression_sample, omega)
  131. #power iterations
  132. for j in np.arange(power_iteration): #same as (XX')^(power iteration)W
  133. B = np.matmul(X_compression_sample, np.matmul(X_compression_sample.T, B))
  134. Q, R = np.linalg.qr(B)
  135. elif axis == 1:
  136. B = np.matmul(omega, X_compression_sample) #l, n_batch
  137. #power iterations
  138. for j in np.arange(power_iteration): #same as H(X'X)^(power iteration)
  139. B = np.matmul(np.matmul(B, X_compression_sample.T), X_compression_sample )
  140. #get QR decomposition
  141. if axis == 0:
  142. Q, R = np.linalg.qr(B)
  143. elif axis == 1:
  144. Q, R = np.linalg.qr(B.T)
  145. else:
  146. print("%s: get_Q_compression: unknown value ( %0.5E ) for axis parameter. Exiting." % (script_name, axis),flush=True)
  147. sys.exit(1)
  148. if VERBOSE_FLAG:
  149. print("%s: get_Q_compression: Q.shape after QR decomposition = " % (script_name), flush=True)
  150. print(Q.shape, flush=True)
  151. print("%s: get_Q_compression: R.shape after QR decomposition = " % (script_name), flush=True)
  152. print(R.shape, flush=True)
  153. return Q
  154. def get_reconstruction_error(X, Q, VERBOSE_FLAG = False ):
  155. """
  156. reconstruction error in terms of compression; how well does compressed X reconstruct back to original full X
  157. """
  158. if VERBOSE_FLAG:
  159. print("%s: get_reconstruction_error: X.shape = " % (script_name), flush=True)
  160. print(X.shape, flush=True)
  161. print("%s: get_reconstruction_error: Q.shape = " % (script_name), flush=True)
  162. print(Q.shape, flush=True)
  163. return np.power(np.linalg.norm(X - np.matmul(Q, np.matmul(Q.T, X))), 2)
  164. ## s-opNMF specific functions
  165. def sample_uniform(X, num_batch, batch_size, max_epoch, axis = 1, WITH_REPLACEMENT = False):
  166. """
  167. inputs
  168. X: input matrix of size [m,n]
  169. axis (default=1): which direction to sample from; default is n in X in [m,n] where n is assumed to be subjects
  170. """
  171. #assign indices from 0 to n (where n is number of subjects, or number of columns of X)
  172. idx_list = np.arange(start = 0, stop = np.shape(X)[axis], step = 1)
  173. idx_sample_full = -1 * np.ones(shape = (batch_size, num_batch, max_epoch)) #since indices are nonnegative, initialize indices to bogus value of negative one so that it is easy to check for mistake when debugging by checking if any value of idx_sample_full is negative
  174. for epoch in np.arange(max_epoch):
  175. if VERBOSE_FLAG:
  176. if np.mod(epoch, 100) == 0: #don't print out all epochs, will print out too many too often
  177. print("%s: sample_uniform: sampling %i-th epoch indices." % (script_name, epoch), flush=True)
  178. idx_sample_full[:,:,epoch] = np.random.choice(idx_list, size = (batch_size, num_batch), replace = WITH_REPLACEMENT)
  179. return idx_sample_full
  180. def sample_dpp(evalue, evector, k):
  181. """
  182. sample a set Y from a dpp. evalue, evector are a decomposed kernel, and k is (optionally) the size of the set to return
  183. :param evalue: eigenvalue
  184. :param evector: normalized eigenvector
  185. :param k: number of cluster
  186. :return:
  187. """
  188. if k == None:
  189. # choose eigenvectors randomly
  190. evalue = np.divide(evalue, (1 + evalue))
  191. v = np.where(np.random.random(evalue.shape[0]) <= evalue)[0]
  192. #evector = np.where(np.random.random(evalue.shape[0]) <= evalue)[0]
  193. else:
  194. v = sample_k(evalue, k) ## v here is a 1d array with size: k
  195. k = v.shape[0]
  196. v = v.astype(int)
  197. v = [i - 1 for i in v.tolist()] ## due to the index difference between matlab & python, here, the element of v is for matlab
  198. V = evector[:, v]
  199. ## iterate
  200. y = np.zeros(k)
  201. for i in range(k, 0, -1):
  202. ## compute probabilities for each item
  203. P = np.sum(np.square(V), axis=1)
  204. P = P / np.sum(P)
  205. # choose a new item to include
  206. y[i-1] = np.where(np.random.rand(1) < np.cumsum(P))[0][0]
  207. y = y.astype(int)
  208. # choose a vector to eliminate
  209. j = np.where(V[y[i-1], :])[0][0]
  210. Vj = V[:, j]
  211. V = np.delete(V, j, 1)
  212. ## Update V
  213. if V.size == 0:
  214. pass
  215. else:
  216. V = np.subtract(V, np.multiply(Vj, (V[y[i-1], :] / Vj[y[i-1]])[:, np.newaxis]).transpose()) ## watch out the dimension here
  217. ## orthogonalize
  218. for m in range(i - 1):
  219. for n in range(m):
  220. V[:, m] = np.subtract(V[:, m], np.matmul(V[:, m].transpose(), V[:, n]) * V[:, n])
  221. V[:, m] = V[:, m] / np.linalg.norm(V[:, m])
  222. y = np.sort(y)
  223. return y
  224. def sample_k(lambda_value, k):
  225. """
  226. Pick k lambdas according to p(S) \propto prod(lambda \in S)
  227. :param lambda_value: the corresponding eigenvalues
  228. :param k: the number of clusters
  229. :return:
  230. """
  231. ## compute elementary symmetric polynomials
  232. E = elem_sym_poly(lambda_value, k)
  233. ## ietrate over the lambda value
  234. num = lambda_value.shape[0]
  235. remaining = k
  236. S = np.zeros(k)
  237. while remaining > 0:
  238. #compute marginal of num given that we choose remaining values from 0:num-1
  239. if num == remaining:
  240. marg = 1
  241. else:
  242. marg = lambda_value[num-1] * E[remaining-1, num-1] / E[remaining, num]
  243. # sample marginal
  244. if np.random.rand(1) < marg:
  245. S[remaining-1] = num
  246. remaining = remaining - 1
  247. num = num - 1
  248. return S
  249. def elem_sym_poly(lambda_value, k):
  250. """
  251. given a vector of lambdas and a maximum size k, determine the value of
  252. the elementary symmetric polynomials:
  253. E(l+1,n+1) = sum_{J \subseteq 1..n,|J| = l} prod_{i \in J} lambda(i)
  254. :param lambda_value: the corresponding eigenvalues
  255. :param k: number of clusters
  256. :return:
  257. """
  258. N = lambda_value.shape[0]
  259. E = np.zeros((k + 1, N + 1))
  260. E[0, :] = 1
  261. for i in range(1, k+1):
  262. for j in range(1, N+1):
  263. E[i, j] = E[i, j - 1] + lambda_value[j-1] * E[i - 1, j - 1]
  264. return E
  265. def sample_dpp_linear(X, num_batch, batch_size, max_epoch, data, axis = 1, print_step = -1, WITH_REPLACEMENT = False, sampling_intermediate_path = "NA"):
  266. import scipy.stats
  267. import scipy.linalg
  268. #assume data is csv of age and sex, is in order of generating the input matrix X
  269. age_sex = data[["age", "sex"]]
  270. age_sex['sex'].replace(['F','M'],[-1,1],inplace=True)
  271. age_sex=age_sex.apply(scipy.stats.zscore) #z-score meta-data
  272. #extract mete-data vectors to create kernel
  273. age = age_sex['age']
  274. age = np.reshape([age],(len(age),1))
  275. age_transpose = np.transpose(age)
  276. sex = age_sex['sex']
  277. sex = np.reshape([sex],(len(sex),1))
  278. sex_transpose = np.transpose(sex)
  279. kern_original = np.dot(X.transpose(), X) #Linear covariance matrix based on imaging features
  280. idx_array_original = np.arange(start = 0, stop = X.shape[axis], step = 1)
  281. # idx_sample_full = -1 * np.ones(shape = (num_batch, batch_size, max_epoch)) #since indices are nonnegative, initialize indices to bogus value of negative one so that it is easy to check for mistake when debugging by checking if any value of idx_sample_full is negative
  282. idx_sample_full = -1 * np.ones(shape = (batch_size, num_batch, max_epoch)) #was originally in wrong shape
  283. finished_epoch = 0
  284. #kernel does not change each epoch; no need to repeat this each epoch; run once and reuse evalue and evector
  285. evalue_original, evector_original = scipy.linalg.eigh(kern_original)
  286. if os.path.isfile(sampling_intermediate_path):
  287. print("sample_dpp_linear: loading sampling indices from sampling_intermediate_path ( %s )." % (sampling_intermediate_path), flush=True)
  288. idx_sample_full = hdf5storage.loadmat(file_name = sampling_intermediate_path, variable_names = ["idx_sample_full"])["idx_sample_full"]
  289. finished_epoch = hdf5storage.loadmat(file_name = sampling_intermediate_path, variable_names = ["finished_epoch"])["finished_epoch"]
  290. print("sample_dpp_linear: loaded sampling indices from sampling_intermediate_path ( %s ). Has indices up to epoch %i" % (sampling_intermediate_path, finished_epoch), flush=True)
  291. if print_step < 0:
  292. print_step = max_epoch #reset to print none to console
  293. for epoch in np.arange(start = finished_epoch, stop = max_epoch, step = 1):
  294. if np.mod(epoch, print_step) == 0:
  295. print("%s: sample_dpp_linear: epoch: %i / %i - generating indices with DPP linear sampling for %i subjects." % (script_name, epoch, max_epoch, X.shape[axis]), flush=True)
  296. #save intermediate
  297. sampling_intermediate_dir = os.path.dirname(sampling_intermediate_path)
  298. if not os.path.isdir(sampling_intermediate_dir):
  299. print("sample_dpp_linear: sampling_intermediate_dir ( %s ) does not exist. Not saving intermediates." % (sampling_intermediate_dir), flush=True)
  300. else:
  301. if np.mod(epoch, 100) == 0:
  302. mdict = {
  303. "idx_sample_full": idx_sample_full,
  304. "finished_epoch": epoch
  305. }
  306. print("sample_dpp_linear: Saving intermediate sampling to ( %s )" % (sampling_intermediate_path), flush=True)
  307. hdf5storage.savemat(file_name = sampling_intermediate_path, mdict=mdict)
  308. specific_sampling_intermediate_basename = os.path.splitext(os.path.basename(sampling_intermediate_path))[0]
  309. specific_sampling_intermediate_path = os.path.join(sampling_intermediate_dir, "%s_epoch%05d.mat" % (specific_sampling_intermediate_basename, epoch))
  310. print("sample_dpp_linear: Saving intermediate sampling to ( %s )" % (specific_sampling_intermediate_path), flush=True)
  311. hdf5storage.savemat(file_name = specific_sampling_intermediate_path, mdict=mdict)
  312. # idx_sample = -1 * np.ones((num_batch,batch_size))
  313. idx_sample = -1 * np.ones(shape = (batch_size, num_batch)) #was originally in wrong shape
  314. idx_array = np.copy(idx_array_original)
  315. kern = np.copy(kern_original)
  316. evector = np.copy(evector_original)
  317. evalue = np.copy(evalue_original)
  318. for sample in range(num_batch):
  319. if VERBOSE_FLAG:
  320. print("%s: sample_dpp_linear: epoch: %i / %i | batch %i / %i - generating indices with sampling." % (script_name, epoch, max_epoch, sample, num_batch), flush=True)
  321. idx = sample_dpp( np.real(evalue),np.real(evector), batch_size)
  322. idx_sample[:, sample] = idx_array[idx]
  323. if not WITH_REPLACEMENT:
  324. kern = np.delete(kern, idx, 0)
  325. kern = np.delete(kern, idx, 1)
  326. idx_array = np.delete(idx_array, idx, 0)
  327. # need to also update evalue and evector
  328. evalue = np.delete(evalue, idx, 0)
  329. evector = np.delete(evector, idx, 0)
  330. evector = np.delete(evector, idx, 1)
  331. idx_sample_full[:,:,epoch] = idx_sample
  332. del idx_sample
  333. #sanity check
  334. if np.sum(idx_sample_full < 0):
  335. print("%s: sample_dpp_linear: error - there are at least %i indices where assignment of index failed to generate valid index greater than or equal to zero." % (script_name, np.sum(idx_sample_full < 0)), flush=True)
  336. return idx_sample_full
  337. def sample_dpp_gaussian(X, num_batch, batch_size, max_epoch, data, sigma = 0.1, axis = 1, print_step = -1, WITH_REPLACEMENT = False):
  338. import scipy.stats
  339. age_sex = data[["age", "sex"]]
  340. age_sex['sex'].replace(['F','M'],[-1,1],inplace=True)
  341. age_sex=age_sex.apply(scipy.stats.zscore) #z-score meta-data
  342. #extract mete-data vectors to create kernel
  343. age = age_sex['age']
  344. age = np.reshape([age],(len(age),1))
  345. age_transpose = np.transpose(age)
  346. sex = age_sex['sex']
  347. sex = np.reshape([sex],(len(sex),1))
  348. sex_transpose = np.transpose(sex)
  349. #kernel
  350. #to-do: make into a function that takes some meta data (optional and required0 and generates a kernel
  351. kern_original = np.exp(-((age-age_transpose)**2+(sex-sex_transpose)**2)/sigma**2) #gaussian
  352. idx_array_original = np.arange(start = 0, stop = X.shape[axis], step = 1)
  353. # idx_sample_full = -1 * np.ones(shape = (num_batch, batch_size, max_epoch)) #since indices are nonnegative, initialize indices to bogus value of negative one so that it is easy to check for mistake when debugging by checking if any value of idx_sample_full is negative
  354. idx_sample_full = -1 * np.ones(shape = (batch_size, num_batch, max_epoch)) #was originally in wrong shape
  355. #kernel does not change each epoch; no need to repeat this each epoch; run once and reuse evalue and evector
  356. evalue_original, evector_original = scipy.linalg.eigh(kern_original)
  357. if print_step < 0:
  358. print_step = max_epoch #reset to print none to console
  359. for epoch in np.arange(max_epoch):
  360. if np.mod(epoch, print_step) == 0:
  361. print("%s: sample_dpp_gaussian: epoch: %i / %i - generating indices with DPP gaussian sampling for %i subjects." % (script_name, epoch, max_epoch, X.shape[axis]), flush=True)
  362. idx_array = np.copy(idx_array_original)
  363. kern = np.copy(kern_original)
  364. evector = np.copy(evector_original)
  365. evalue = np.copy(evalue_original)
  366. # idx_sample = -1 * np.ones((num_batch,batch_size))
  367. idx_sample = -1 * np.ones(shape = (batch_size, num_batch)) #was originally in wrong shape
  368. for sample in range(num_batch):
  369. if VERBOSE_FLAG:
  370. print("%s: sample_dpp_gaussian: epoch: %i / %i | batch %i / %i - generating indices with sampling." % (script_name, epoch, max_epoch, sample, num_batch), flush=True)
  371. idx = sample_dpp(np.real(evalue),np.real(evector), batch_size)
  372. # idx_sample[sample,:] = idx_array[idx]
  373. idx_sample[:,sample] = idx_array[idx]
  374. if not WITH_REPLACEMENT:
  375. kern = np.delete(kern, idx, 0)
  376. kern = np.delete(kern, idx, 1)
  377. idx_array = np.delete(idx_array, idx, 0)
  378. # need to also update evalue and evector
  379. evalue = np.delete(evalue, idx, 0)
  380. evector = np.delete(evector, idx, 0)
  381. evector = np.delete(evector, idx, 1)
  382. idx_sample_full[:,:,epoch] = idx_sample
  383. del idx_sample
  384. #sanity check
  385. if np.sum(idx_sample_full < 0):
  386. print("%s: sample_dpp_linear: error - there are at least %i indices where assignment of index failed to generate valid index greater than or equal to zero." % (script_name, np.sum(idx_sample_full < 0)), flush=True)
  387. return idx_sample_full
  388. def calculate_batch_loss_sum(X, idx_sample, W, SQUARED = True):
  389. """
  390. given indices of subjects for a specific batch, the final W from all the iterations/batches of an epoch, and X for original data to sample the indices from, calculate the sum of batch loss
  391. sum_{i=0}^{last batch}(|X_batch - W_final (W_final * X_batch)|_F^2)
  392. """
  393. batch_loss_sum = 0.0
  394. num_batch = idx_sample.shape[1]
  395. for batch_index in np.arange(num_batch):
  396. idx_sample_selected = (idx_sample[:, int(batch_index)]).astype(int)
  397. X_sampled = X[:, idx_sample_selected]
  398. del idx_sample_selected
  399. #for the sake of consistency, let us call the calculate_error function from utils module rather than calling frobenius norm function from scratch
  400. batch_loss_sum = batch_loss_sum + utils.calculate_error(X = X_sampled, W = W, H = np.matmul(W.T, X_sampled), SQUARED = SQUARED)
  401. del X_sampled
  402. return batch_loss_sum
  403. ## argparser
  404. parser=argparse.ArgumentParser(
  405. description = "This script is a mostly python conversion of opnmf.m matlab script. It runs orthonormal projective non-negative matrix factorization. In addition, it has options to allow SVD or QR decomposition on input data matrix to optimize update rules.",
  406. epilog = "Written by Sung Min Ha ([email hidden])",
  407. add_help=True,
  408. )
  409. parser.add_argument("-i", "--inputFile",
  410. type = str,
  411. dest = "input_path",
  412. required = False,
  413. default = feret_path,
  414. help = "Full path to the mat file containing variable X that contains input data to work with or a csv file containing list of nii.gz or mgh files to read and construct into input matrix."
  415. )
  416. parser.add_argument("-d", "--demographic_data_path",
  417. type = str,
  418. dest = "demographic_data_path",
  419. required = False,
  420. default = csv_path,
  421. help = "Full path to the csv file in the same order as the input X file of subject domain with relevant columns age and sex for sampling methods (such as dpp)."
  422. )
  423. parser.add_argument("-k", "--targetRank",
  424. type = int,
  425. dest = "target_rank",
  426. required = False,
  427. default = 40,
  428. help = "What is the rank of the NMF you intend to run? This is the number of components that will be generated."
  429. )
  430. parser.add_argument("-m", "--maxEpoch",
  431. type = int,
  432. dest = "max_epoch",
  433. required = False,
  434. default = 5.0e4,
  435. help = "Max number of iterations to optimize over. Note that if other stopping criteria are achived (e.g. tolerance), then the algorithm may stop before reaching this max number of iterations."
  436. )
  437. parser.add_argument("-t", "--tol",
  438. type = float,
  439. dest = "tol",
  440. required = False,
  441. default = 0.0e0,
  442. help = "Tolerance value to use as threshold for stopping criterion. If the diffW = norm(W-W_old) / norm(W) is less than this tolerance value, then the iterations would stop for optimization regardless of whether max_epoch has been reached, under the assumption that the cost function has stabilized (plateau) and has reached close to local minimum."
  443. )
  444. parser.add_argument("-o", "--outputParentDir",
  445. type = str,
  446. dest = "output_parent_dir",
  447. required = False,
  448. default = os.path.join("/scratch/%s" % (username), "output_directory"),
  449. help = "Path to output mat file that contain the outputs."
  450. )
  451. parser.add_argument("-0", "--initMeth",
  452. type = str,
  453. dest = "init_method",
  454. required = False,
  455. default = "nndsvd",
  456. help = "Method for initializing w0 for component. (random, nndsvd, nndsvda, nndsvdar)"
  457. )
  458. parser.add_argument("-u", "--updateMeth",
  459. type = str,
  460. dest = "update_meth",
  461. required = False,
  462. default = "mem",
  463. help = "Method for update W. (mem, original)"
  464. )
  465. parser.add_argument("--multiplicativeUpdateMeth",
  466. type = str,
  467. dest = "multiplicative_update_method",
  468. required = False,
  469. default = "normalize",
  470. help = "Method for multiplicative update of W. (normalize, constant, adaptive, adaptiveNormalize, quadratic, quadraticOrthonormal)"
  471. )
  472. parser.add_argument("-s", "--samplingMeth",
  473. type = str,
  474. dest = "sampling_method",
  475. required = False,
  476. default = "uniform",
  477. help = "Method for sampling of batch. (uniform, dppgaussian, dpplinear)"
  478. )
  479. parser.add_argument("-z", "--sampleSize",
  480. type = int,
  481. dest = "batch_size",
  482. required = False,
  483. default = 100,
  484. help = "Number of subjects per batch"
  485. )
  486. parser.add_argument("--printStep",
  487. type = int,
  488. dest = "print_step",
  489. required = False,
  490. default = 1.0e1,
  491. help = "Print progress every this many epochs."
  492. )
  493. parser.add_argument("--saveStep",
  494. type = int,
  495. dest = "save_step",
  496. required = False,
  497. default = 1.0e3,
  498. help = "save progress (batch wise error per iteration, full data wise error, sparsity, elapsed time) every this many epochs into intermediate save files so that you can use those to track and plot changes over iterations/batches."
  499. )
  500. parser.add_argument("--restartStep",
  501. type = int,
  502. dest = "restart_step",
  503. required = False,
  504. default = 1.0e0,
  505. help = "save progress every this many epochs. This will maintain a single file that gets overwritten this many epochs to serve a restart/checkpoint if the code fails."
  506. )
  507. parser.add_argument("--rho",
  508. type = float,
  509. dest = "rho",
  510. required = False,
  511. default = 0.25,
  512. help = "For constant power or adaptive multiplicative update."
  513. )
  514. parser.add_argument("--eta",
  515. type = float,
  516. dest = "eta",
  517. required = False,
  518. default = 0.1,
  519. help = "For adaptive multiplicative update."
  520. )
  521. parser.add_argument("--sigma",
  522. type = float,
  523. dest = "sigma",
  524. required = False,
  525. default = 0.1,
  526. help = "For dpp gaussian kernel generation."
  527. )
  528. parser.add_argument("-V", "--verbose",
  529. action = 'store_true',
  530. dest = "VERBOSE_FLAG",
  531. help = "Extra printouts for debugging."
  532. )
  533. parser.add_argument("-D", "--debug",
  534. action = 'store_true',
  535. dest = "DEBUG_FLAG",
  536. help = "EXTRA EXTRA printouts for debugging."
  537. )
  538. parser.add_argument("--withReplacement",
  539. action = 'store_true',
  540. dest = "WITH_REPLACEMENT_FLAG",
  541. help = "If set to true by calling this flag on, DPP sampling will happen with replacement instead of removing already selected sample"
  542. )
  543. parser.add_argument("--QCompressionBatchSize",
  544. type = int,
  545. dest = "Q_X_compression_size",
  546. required = False,
  547. default = 1.0e3,
  548. help = "When calculating Q [n,l] for compression data in subject dimension, we can either calculate Q on full X (preferable scenario if X is within reasonable size limit), or a mini-batch of X (if X is too large to calculate QR decomposition on)."
  549. )
  550. parser.add_argument("--oversampling",
  551. type = int,
  552. dest = "oversampling",
  553. required = False,
  554. default = 10,
  555. help = "When performing compression, we want l = k + k_ov, where k is the target rank. This is the value of k_ov in the equation. Note that if k + k_ov > n, then l will be set to n. If k + k_ov < k + 10, then l will reset to k + 10"
  556. )
  557. parser.add_argument("--powerIteration",
  558. type = int,
  559. dest = "power_iteration",
  560. required = False,
  561. default = 4,
  562. help = "What is the power iteration for compression matrix Q generation? Defaults to 4."
  563. )
  564. parser.add_argument("--compressionPerBatch",
  565. action = 'store_true',
  566. dest = "COMPRESSION_PER_BATCH_FLAG",
  567. help = "Instead of calculating compression matrix Q from full X before the multiplicative updates and reusing it over and over again, calculate Q PER each mini batch of X (X_p) sampled from X."
  568. )
  569. parser.add_argument("--notOrthonormal",
  570. action = 'store_false',
  571. dest = "ORTHONORMAL_FLAG",
  572. help = "Use the orthonormal projective update instead of projective update"
  573. )
  574. parser.add_argument("--speedOptimized",
  575. action = 'store_false',
  576. dest = "MEM_FLAG",
  577. help = "Do you want to use speed optimized but memory inefficient multiplicative update where you store XX^T in the memory and reuse it."
  578. )
  579. parser.add_argument("--noZeroRemoval",
  580. action = 'store_false',
  581. dest = "ZERO_REMOVAL_FLAG",
  582. help = "turn off the flag to remove all rows in the X input matrix where it is zero across the entire row."
  583. )
  584. parser.add_argument("--noSmallValueReset",
  585. action = 'store_false',
  586. dest = "SMALL_VALUE_RESET_FLAG",
  587. help = "turn off reset by 1.0e-16 (or the value set by min_reset_value variable) any values in W smaller 1.0e-16 (or the value set by min_reset_value variable)?"
  588. )
  589. parser.add_argument("--noSmallValueResetInit",
  590. action = 'store_false',
  591. dest = "SMALL_VALUE_RESET_INIT_FLAG",
  592. help = "turn off reset by 1.0e-16 (or the value set by min_reset_value variable) any values in W smaller 1.0e-16 (or manually set min_reset_value variable) before starting for loop?"
  593. )
  594. parser.add_argument("--minResetValue",
  595. type = str,
  596. dest = "min_reset_value",
  597. default = "1.0e-16",
  598. help = "reset by this value every iteration for values of W less than this value? The default value is set to be the same as in https://github.com/asotiras/brainparts/blob/master/opnmf_mem.m: 1.0e-16. Value taken in as string, then converted to float internally."
  599. )
  600. parser.add_argument("--epsilon",
  601. type = str,
  602. dest = "EPSILON",
  603. default = str(np.finfo(np.float64).eps),
  604. help = "what value to use in demoninator of update rule to avoid division by zero. Default value is the minimum float64 (double) value possible on your system. The default value of np.finfo(np.float64).eps approximately equal to 2.2e-16 matches the value of eps(1) in MatLab. Value taken in as string, then converted to float internally."
  605. )
  606. parser.add_argument("--iterationPerBatch",
  607. type = int,
  608. dest = "iteration_per_batch",
  609. default = int(1),
  610. help = "How many iterations of multiplicative updates do you want to perform per batch per epoch? default is 1."
  611. )
  612. #parse argparser
  613. args = parser.parse_args()
  614. input_path = args.input_path
  615. demographic_data_path = args.demographic_data_path
  616. target_rank = args.target_rank
  617. tol = args.tol #tolerance
  618. output_parent_dir = args.output_parent_dir
  619. init_method = args.init_method
  620. sampling_method = args.sampling_method
  621. multiplicative_update_method = args.multiplicative_update_method
  622. batch_size = args.batch_size
  623. update_meth = args.update_meth
  624. print_step = args.print_step
  625. save_step = args.save_step
  626. restart_step = args.restart_step
  627. VERBOSE_FLAG = args.VERBOSE_FLAG
  628. DEBUG_FLAG = args.DEBUG_FLAG #also save all intermediates at save point
  629. WITH_REPLACEMENT_FLAG = args.WITH_REPLACEMENT_FLAG
  630. # iter_power = args.iter_power
  631. max_epoch = int(args.max_epoch) #same as number of iterations for opNMF, e.g.) 50K
  632. rho = args.rho #1/4 by default
  633. eta = args.eta #0.1 by default
  634. sigma = args.sigma #0.1 by default
  635. SQUARED_ERROR_FLAG = True
  636. iteration_per_batch = int(args.iteration_per_batch)
  637. reset_small_value_threshold = 1.0e-16
  638. Q_X_compression_size = args.Q_X_compression_size
  639. oversampling = args.oversampling
  640. power_iteration = args.power_iteration
  641. COMPRESSION_PER_BATCH_FLAG = args.COMPRESSION_PER_BATCH_FLAG
  642. #update related parameters
  643. MEM_FLAG = args.MEM_FLAG
  644. ORTHONORMAL_FLAG = args.ORTHONORMAL_FLAG
  645. #additional resetting/small value/zero handling parameters
  646. SMALL_VALUE_RESET_FLAG = args.SMALL_VALUE_RESET_FLAG
  647. SMALL_VALUE_RESET_INIT_FLAG = args.SMALL_VALUE_RESET_INIT_FLAG
  648. min_reset_value = float(fractions.Fraction(str(args.min_reset_value)))
  649. if SMALL_VALUE_RESET_INIT_FLAG and not SMALL_VALUE_RESET_FLAG:
  650. SMALL_VALUE_RESET_FLAG = True #if you are initializing to reset small value before for loop, you should do so for inside for loop as well
  651. utils.print_verbose("Setting SMALL_VALUE_RESET_FLAG to %r since SMALL_VALUE_RESET_INIT_FLAG is %r." % (SMALL_VALUE_RESET_FLAG, SMALL_VALUE_RESET_INIT_FLAG), script_name = script_name )
  652. EPSILON = float(fractions.Fraction(str(args.EPSILON)))
  653. ZERO_REMOVAL_FLAG = args.ZERO_REMOVAL_FLAG
  654. if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
  655. #store original rho
  656. rho0 = rho
  657. ORTHONORMAL_FLAG = True
  658. #sanity check: check whether correct configuration is setup for speed vs mem
  659. MEM_FLAG = False
  660. if (update_meth == "mem"):
  661. pass
  662. MEM_FLAG = True
  663. elif (update_meth == "original"):
  664. pass
  665. else:
  666. utils.print_flush("ERROR - Unknown update_meth (%s). Accepted values are mem and original. Exiting." % (update_meth), script_name = script_name)
  667. sys.exit(1)
  668. #sanity check: turn verbose flag on if debug flag is on
  669. if DEBUG_FLAG:
  670. VERBOSE_FLAG = True
  671. #print out final variables after argparsing
  672. print("\n\n", flush=True)
  673. utils.print_flush(string_variable = "-----Variables-----\n\n", script_name = script_name )
  674. utils.print_flush(string_variable = "input_path: ( %s )" % (input_path), script_name = script_name)
  675. utils.exit_if_not_exist_file(input_path)
  676. utils.print_flush(string_variable = "target_rank: ( %i )" % (target_rank), script_name = script_name)
  677. utils.print_flush(string_variable = "tol: ( %0.5E )" % (tol), script_name = script_name)
  678. utils.print_flush(string_variable = "output_parent_dir: ( %s )" % (output_parent_dir), script_name = script_name)
  679. utils.print_flush(string_variable = "init_method: ( %s )" % (init_method), script_name = script_name)
  680. utils.print_flush(string_variable = "sampling_method: ( %s )" % (sampling_method), script_name = script_name)
  681. utils.print_flush(string_variable = "multiplicative_update_method: ( %s )" % (multiplicative_update_method), script_name = script_name)
  682. utils.print_flush(string_variable = "ORTHONORMAL_FLAG: ( %r )" % (ORTHONORMAL_FLAG), script_name = script_name)
  683. utils.print_flush(string_variable = "MEM_FLAG: ( %r )" % (MEM_FLAG), script_name = script_name)
  684. utils.print_flush(string_variable = "update_meth: ( %s )" % (update_meth), script_name = script_name)
  685. utils.print_flush(string_variable = "print_step: ( %i )" % (print_step), script_name = script_name)
  686. utils.print_flush(string_variable = "restart_step: ( %i )" % (restart_step), script_name = script_name)
  687. utils.print_flush(string_variable = "save_step: ( %i )" % (save_step), script_name = script_name)
  688. utils.print_flush(string_variable = "max_epoch: ( %i )" % (max_epoch), script_name = script_name)
  689. utils.print_flush(string_variable = "rho: ( %0.5E )" % (rho), script_name = script_name)
  690. if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
  691. utils.print_flush(string_variable = "rho0: ( %0.5E )" % (rho0), script_name = script_name)
  692. if sampling_method == "dppgaussian":
  693. utils.print_flush(string_variable = "sigma: ( %0.5E )" % (sigma), script_name = script_name)
  694. utils.print_flush(string_variable = "eta: ( %0.5E )" % (eta), script_name = script_name)
  695. #compression related values
  696. utils.print_flush(string_variable = "batch_size: ( %i )" % (batch_size), script_name = script_name)
  697. utils.print_flush(string_variable = "WITH_REPLACEMENT_FLAG: ( %r )" % (WITH_REPLACEMENT_FLAG), script_name = script_name)
  698. utils.print_flush(string_variable = "Q_X_compression_size: ( %i )" % (Q_X_compression_size), script_name = script_name)
  699. utils.print_flush(string_variable = "oversampling: ( %i )" % (oversampling), script_name = script_name)
  700. utils.print_flush(string_variable = "power_iteration: ( %i )" % (power_iteration), script_name = script_name)
  701. utils.print_flush(string_variable = "COMPRESSION_PER_BATCH_FLAG: ( %r )" % (COMPRESSION_PER_BATCH_FLAG), script_name = script_name)
  702. utils.print_flush(string_variable = "EPSILON: ( %0.2E )" % (EPSILON), script_name = script_name)
  703. utils.print_verbose(string_variable = "EPSILON precision: %i" % (utils.get_numpy_precision(EPSILON)), script_name = script_name, VERBOSE_FLAG = VERBOSE_FLAG)
  704. utils.print_flush(string_variable = "SMALL_VALUE_RESET_INIT_FLAG: ( %r )" %
  705. (SMALL_VALUE_RESET_INIT_FLAG), script_name = script_name)
  706. utils.print_verbose(string_variable = "SMALL_VALUE_RESET_INIT_FLAG ( %r ): reset elements of W < %0.2E to %0.2E prior to the start of the for loop" % (SMALL_VALUE_RESET_INIT_FLAG, min_reset_value, min_reset_value), script_name = script_name, VERBOSE_FLAG = VERBOSE_FLAG)
  707. utils.print_flush(string_variable = "SMALL_VALUE_RESET_FLAG: ( %r )" % (SMALL_VALUE_RESET_FLAG), script_name = script_name)
  708. utils.print_verbose(string_variable = "SMALL_VALUE_RESET_FLAG ( %r ): reset elements of W < %0.2E to %0.2E during each of the for loop" % (SMALL_VALUE_RESET_FLAG, min_reset_value, min_reset_value), script_name = script_name, VERBOSE_FLAG = VERBOSE_FLAG)
  709. utils.print_flush(string_variable = "min_reset_value: ( %0.2E )" % (min_reset_value), script_name = script_name)
  710. utils.print_flush(string_variable = "ZERO_REMOVAL_FLAG: ( %r )" % (ZERO_REMOVAL_FLAG), script_name = script_name)
  711. utils.print_flush(string_variable = "VERBOSE_FLAG: ( %r )" % (VERBOSE_FLAG), script_name = script_name)
  712. utils.print_flush(string_variable = "DEBUG_FLAG: ( %r )" % (DEBUG_FLAG), script_name = script_name)
  713. if ORTHONORMAL_FLAG:
  714. nmf_method = "opNMF"
  715. else:
  716. nmf_method = "pNMF"
  717. nmf_method = "cs-%s" % (nmf_method) #cs-opNMF or cs-pNMF
  718. ## initialize elapse time counts
  719. total_elapsed_time = 0.0 #all time elapsed from this point on
  720. initialization_elapsed_time = 0.0 #for initializations
  721. saving_elapsed_time = 0.0 #for intermediate savings and loading from restart
  722. statistics_elapsed_time = 0.0 #for calculating statistics (error/sparsity on the fly)
  723. update_elapsed_time = 0.0 #for multiplicative update
  724. sampling_elapsed_time = 0.0 #for sampling batches and selecting those batches as indices
  725. compression_elapsed_time = 0.0 #for compression (calculation and applying of Q matrices)
  726. total_start_time = time.time()
  727. total_start_time_string = datetime.datetime.fromtimestamp(total_start_time).strftime('%Y-%m-%d %H:%M:%S')
  728. utils.print_flush("total start time\t: %s" % (total_start_time_string), script_name = script_name)
  729. ## Input Data Loading
  730. saving_start_time = time.time()
  731. utils.exit_if_not_exist_file(file_path = input_path, script_name = script_name)
  732. utils.print_flush("Loading input variable X from ( %s )." % (input_path), script_name = script_name)
  733. X = utils.load_hdf5storage_data(file_path = input_path, variable_name = "X")
  734. m = np.shape(X)[0]
  735. n = np.shape(X)[1]
  736. if np.shape(np.shape(X))[0] != 2:
  737. utils.print_flush("ERROR: X loaded from ( %s ) is not 2-D, but %i-D. Exiting." % (input_path, np.shape(np.shape(X))[0]), script_name = script_name)
  738. sys.exit(1)
  739. utils.print_verbose("Loaded input variable X of shape [%i, %i] from ( %s )." % (m, n, input_path), script_name = script_name)
  740. #calculate l
  741. l = np.amin([n, np.amax([target_rank + oversampling, target_rank + 10])])
  742. utils.print_flush(string_variable = "l: ( %i )" % (l), script_name = script_name)
  743. outdir = restart.get_output_directory(
  744. output_parent_dir = output_parent_dir,
  745. nmf_method = nmf_method,
  746. target_rank = target_rank,
  747. tol = tol,
  748. max_iter = max_epoch,
  749. update_method = multiplicative_update_method,
  750. ORTHONORMAL_FLAG = ORTHONORMAL_FLAG,
  751. init_method = init_method,
  752. MEM_FLAG = MEM_FLAG,
  753. sampling_method = sampling_method,
  754. rho0 = rho, #since rho0 is initialized to rho if adaptive and not intialized at all when not adaptive update, can simply put rho here in place of rho0
  755. eta = eta,
  756. batch_size = batch_size,
  757. X_compression_size =Q_X_compression_size,
  758. l = l,
  759. ZERO_REMOVAL_FLAG = ZERO_REMOVAL_FLAG,
  760. SMALL_VALUE_RESET_INIT_FLAG = SMALL_VALUE_RESET_INIT_FLAG,
  761. SMALL_VALUE_RESET_FLAG = SMALL_VALUE_RESET_FLAG,
  762. min_reset_value = min_reset_value,
  763. EPSILON = EPSILON,
  764. iteration_per_batch = iteration_per_batch,
  765. power_iteration_compression = power_iteration,
  766. COMPRESSION_PER_BATCH_FLAG = COMPRESSION_PER_BATCH_FLAG,
  767. script_name = script_name
  768. )
  769. output_basename_prefix = "%s_%s" % (nmf_method, update_meth)
  770. #define output directory
  771. #To-Do: replace with restart.py output directory function
  772. # outdir = os.path.join(output_parent_dir, "cs-opNMF", "targetRank%i" % (target_rank), "init%s" % (init_method), "update%s" % (multiplicative_update_method), "tol%0.2E" % (tol), "maxEpoch%0.2E" % (max_epoch), "sampling%s" % (sampling_method), "batchSize%i" % (batch_size))
  773. # if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize" or multiplicative_update_method == "quadratic" or multiplicative_update_method == "quadraticOrthonormal":
  774. # outdir = os.path.join(outdir, "rhoInit%0.2E" % (rho0))
  775. # if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
  776. # outdir = os.path.join(outdir, "eta%0.2E" % (eta))
  777. # if sampling_method == "dppggaussian":
  778. # outdir = os.path.join(outdir, "sigma%0.2E" % (sigma))
  779. # output_path = os.path.join(outdir, "%s.mat" % (output_basename_prefix))
  780. output_path = restart.get_output_path(
  781. output_parent_dir = output_parent_dir,
  782. nmf_method = nmf_method,
  783. target_rank = target_rank,
  784. tol = tol,
  785. max_iter = max_epoch,
  786. update_method = multiplicative_update_method,
  787. ORTHONORMAL_FLAG = ORTHONORMAL_FLAG,
  788. init_method = init_method,
  789. MEM_FLAG = MEM_FLAG,
  790. basename_prefix = output_basename_prefix,
  791. sampling_method = sampling_method,
  792. rho0 = rho, #since rho0 is initialized to rho if adaptive and not intialized at all when not adaptive update, can simply put rho here in place of rho0
  793. eta = eta,
  794. batch_size = batch_size,
  795. X_compression_size =Q_X_compression_size,
  796. l = l,
  797. ZERO_REMOVAL_FLAG = ZERO_REMOVAL_FLAG,
  798. SMALL_VALUE_RESET_INIT_FLAG = SMALL_VALUE_RESET_INIT_FLAG,
  799. SMALL_VALUE_RESET_FLAG = SMALL_VALUE_RESET_FLAG,
  800. min_reset_value = min_reset_value,
  801. EPSILON = EPSILON,
  802. iteration_per_batch = iteration_per_batch,
  803. power_iteration_compression = power_iteration,
  804. COMPRESSION_PER_BATCH_FLAG = COMPRESSION_PER_BATCH_FLAG,
  805. script_name = script_name
  806. )
  807. utils.print_flush(string_variable = "outdir: ( %s )" % (outdir), script_name = script_name)
  808. utils.print_flush(string_variable = "output_path: ( %s )" % (output_path), script_name = script_name)
  809. #output path
  810. restart_path = os.path.join(outdir, "%s_restart.mat" % (output_basename_prefix))
  811. restart_path_old = os.path.join(outdir, "%s_restart_old.mat" % (output_basename_prefix))
  812. initialization_path = os.path.join(outdir, "%s_initialization.mat" % (output_basename_prefix))
  813. sampling_intermediate_path = os.path.join(outdir, "%s_sampling.mat" % (output_basename_prefix))
  814. utils.print_flush(string_variable = "restart_path: ( %s )" % (restart_path), script_name = script_name)
  815. utils.print_flush(string_variable = "initialization_path: ( %s )" % (initialization_path), script_name = script_name)
  816. utils.print_flush(string_variable = "sampling_intermediate_path: ( %s )" % (sampling_intermediate_path), script_name = script_name)
  817. #sanity check: does parent output directory exist?
  818. utils.print_verbose("Checking if output directory ( %s ) exists." % (output_parent_dir), script_name = script_name, VERBOSE_FLAG = VERBOSE_FLAG)
  819. utils.exit_if_not_exist_dir(dir_path = output_parent_dir, script_name = script_name)
  820. # outdir = os.path.join(outdir, "Q_X_compression_size%i" % (Q_X_compression_size), "l%i" % (l), "powerIter%0.2E" % (power_iteration))
  821. # utils.print_flush(string_variable = "outdir: ( %s )" % (outdir), script_name = script_name)
  822. #sanity check: does output exist?
  823. if not os.path.isdir(outdir):
  824. os.makedirs(outdir)
  825. utils.exit_if_exist_file(file_path = output_path, script_name = script_name)
  826. #load demographics data if not uniform sampling
  827. if (sampling_method == "dpplinear") or (sampling_method == "dppgaussian"):
  828. utils.print_flush(string_variable = "demographic_data_path: ( %s )" %(demographic_data_path), script_name = script_name)
  829. utils.exit_if_not_exist_file(demographic_data_path)
  830. demographic_data = pd.read_csv(demographic_data_path)
  831. #sanity check: does k <= n and k <= m
  832. utils.print_verbose("sanity check on dimensions of data matrix X ( [ %i, %i ] ) to target rank k = %i." % (m, n, target_rank), script_name = script_name, VERBOSE_FLAG = VERBOSE_FLAG)
  833. if target_rank > m:
  834. utils.print_flush("target_rank (%i) > m (%i, the number of features or rows of X). target_rank must be less than m. Exiting." % (target_rank, m), script_name = script_name)
  835. sys.exit(1)
  836. if target_rank > n:
  837. utils.print_flush("target_rank (%i) > n (%i, the number of subjects or columns of X). target_rank must be less than n. Exiting." % (target_rank, n), script_name = script_name)
  838. sys.exit(1)
  839. saving_elapsed_time = saving_elapsed_time + (time.time() - saving_start_time)
  840. del saving_start_time
  841. ## Initializations
  842. #To-Do: load from initialization file if it exists instead of always initializing from scratch
  843. if os.path.isfile(initialization_path):
  844. saving_start_time = time.time()
  845. utils.print_verbose("Loading intialization data from ( %s )." % (initialization_path), script_name = script_name, VERBOSE_FLAG = VERBOSE_FLAG)
  846. w0 = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "w0", script_name = script_name)
  847. h0 = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "w0", script_name = script_name)
  848. saving_elapsed_time = saving_elapsed_time + ( time.time() - saving_start_time )
  849. del saving_start_time
  850. else:
  851. initialization_start_time = time.time()
  852. utils.print_verbose("Initializing W and H with %s intialization." % (init_method), script_name = script_name, VERBOSE_FLAG = VERBOSE_FLAG)
  853. if (init_method == "nndsvd") or (init_method == "nndsvda") or (init_method == "nndsvdar") or (init_method == "random"):
  854. w0, h0 = initialize_nmf._initialize_nmf(X, target_rank, init=init_method, eps=1e-6, random_state=None)
  855. else:
  856. utils.print_flush("ERROR - Unknown init_method (%s). Exiting." % (init_method), script_name = script_name)
  857. initialization_elapsed_time = initialization_elapsed_time + (time.time() - initialization_start_time)
  858. del initialization_start_time
  859. #initialize rest of variables (XX, W_old, diffW, etc.)
  860. initialization_start_time = time.time()
  861. #XX^T if not using MEM mode
  862. utils.print_verbose("Calculating XX^T to store in memory.", script_name = script_name)
  863. if (MEM_FLAG == False):
  864. XX = np.matmul(X, np.transpose(X))
  865. #initialize W to w0
  866. W = w0
  867. W_old = W
  868. #initialize diffW for print purposes and for stopping criterion
  869. diffW = 0.0
  870. initialization_elapsed_time = initialization_elapsed_time + (time.time() - initialization_start_time)
  871. del initialization_start_time
  872. ## Sampling
  873. #load idx_sample if inititialization path exists
  874. if os.path.isfile(initialization_path):
  875. saving_start_time = time.time()
  876. num_batch = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "num_batch", script_name = script_name)
  877. idx_sample = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "idx_sample", script_name = script_name)
  878. saving_elapsed_time = saving_elapsed_time + ( time.time() - saving_start_time )
  879. del saving_start_time
  880. #if initialization path does not exist, sample from scratch
  881. else:
  882. sampling_start_time = time.time()
  883. utils.print_verbose("Calculating num_batch based on n ( %i ) and batch_size ( %i )" % (n, batch_size), script_name = script_name)
  884. num_batch = int(np.ceil(np.divide(float(n), float(batch_size)))) #we need this value as int
  885. if (num_batch != np.divide(float(n), float(batch_size))):
  886. utils.print_flush("ERROR: currently, if your n (%i) divided by batch_size (%i) = (%f) does not end up as an integer (%i), then this code fails." % (n, batch_size, np.divide(float(n), float(batch_size)), num_batch), script_name = script_name)
  887. #generate indices for all iterations/epochs
  888. utils.print_flush("Sampling indices of subjects using %s sampling method" % (sampling_method), script_name = script_name)
  889. if (batch_size == n):
  890. print("%s: Since the number of subjects (X.shape[1]) is ( %i ) and equal to size of a batch ( %i ), no shuffling will be performed." % (script_name, n, batch_size), flush=True)
  891. idx_sample = -1 * np.ones(shape = (n, 1, max_epoch))
  892. for epoch in np.arange(max_epoch):
  893. idx_sample[:,:,epoch] = np.reshape(np.arange(n), newshape = (n, 1))
  894. print("%s: idx_sample = " % (script_name), flush=True)
  895. print(idx_sample, flush=True)
  896. elif (sampling_method == "uniform"):
  897. idx_sample = sample_uniform(X = X, num_batch = int(num_batch), batch_size = int(batch_size), max_epoch = int(max_epoch), WITH_REPLACEMENT=WITH_REPLACEMENT_FLAG)
  898. elif (sampling_method == "dpplinear"):
  899. #save every 100 epochs
  900. idx_sample = sample_dpp_linear(X = X, num_batch = num_batch, batch_size = batch_size, max_epoch = max_epoch, data = demographic_data, axis = 1, print_step = int(np.round(max_epoch / 100.0)), WITH_REPLACEMENT=WITH_REPLACEMENT_FLAG, sampling_intermediate_path = sampling_intermediate_path )
  901. elif (sampling_method == "dppgaussian"):
  902. idx_sample = sample_dpp_gaussian(X = X, num_batch = num_batch, batch_size = batch_size, max_epoch = max_epoch, data = demographic_data, axis = 1, print_step = int(np.round(max_epoch / 100.0)), sigma = sigma, WITH_REPLACEMENT=WITH_REPLACEMENT_FLAG )
  903. else:
  904. utils.print_flush("Unknown sampling method (%s). Exiting.", script_name = script_name)
  905. sys.exit(1)
  906. utils.print_flush("Finished sampling indices of subjects using %s sampling method" % (sampling_method), script_name = script_name)
  907. sampling_elapsed_time = sampling_elapsed_time + (time.time() - sampling_start_time)
  908. del sampling_start_time
  909. #calculate Q for compression
  910. if not COMPRESSION_PER_BATCH_FLAG:
  911. if os.path.isfile(initialization_path):
  912. saving_start_time = time.time()
  913. Q = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "Q", script_name = script_name)
  914. saving_elapsed_time = saving_elapsed_time + ( time.time() - saving_start_time )
  915. del saving_start_time
  916. else:
  917. compression_start_time = time.time()
  918. Q_full = get_Q_compression(X_hat = X, X_compression_size = Q_X_compression_size, power_iteration = power_iteration, l = l, axis = 1, VERBOSE_FLAG = VERBOSE_FLAG)
  919. if VERBOSE_FLAG:
  920. print("Q.shape = ", flush=True)
  921. print(Q.shape, flush=True) #should be l by l?
  922. compression_elapsed_time = compression_elapsed_time + (time.time() - compression_start_time)
  923. del compression_start_time
  924. #sanity check: if Q_X_compression_size > batch_size, select first batch_size of Q_X_compression_size of Q to use for compression
  925. compression_start_time = time.time()
  926. if Q_X_compression_size < batch_size:
  927. print("%s: ERROR - you generated Q compression matrix from %i subjects of X_hat but you are using mini-batches of size %i subjects, where %i > %i. Exiting." % (script_name, Q_X_compression_size, batch_size, Q_X_compression_size, batch_size), flush=True)
  928. sys.exit(1)
  929. elif Q_X_compression_size == batch_size:
  930. Q = Q_full.copy()
  931. else:
  932. Q = Q_full[0:batch_size, :]
  933. if VERBOSE_FLAG:
  934. print("%s: Extracting Q (for actually compressing X_tilde mini-batch repeatedly) of shape [%i, %i] out of Q_full (from X_hat for generating Q) of shape [%i, %i]" % (script_name, Q.shape[0], Q.shape[1], Q_full.shape[0], Q_full.shape[1]), flush=True)
  935. del Q_full
  936. compression_elapsed_time = compression_elapsed_time + (time.time() - compression_start_time)
  937. del compression_start_time
  938. utils.print_flush("COMPRESSION_PER_BATCH_FLAG: ( %r ) \t| Q.shape = [ %i, %i ]" % (COMPRESSION_PER_BATCH_FLAG, Q.shape[0], Q.shape[1]), script_name = script_name)
  939. utils.print_verbose("X.shape = [%i, %i]" % (m, n) , script_name = script_name)
  940. utils.print_verbose("X.min max = [%0.15f, %0.15f]" % (np.amin(X, axis = None), np.amax(X, axis = None) ) , script_name = script_name)
  941. utils.print_verbose("w0.shape = [%i, %i]" % (np.shape(w0)[0], np.shape(w0)[1]) , script_name = script_name)
  942. utils.print_verbose("w0.min max = [%0.15f, %0.15f]" % (np.amin(w0, axis = None), np.amax(w0, axis = None) ) , script_name = script_name)
  943. utils.print_verbose("h0.shape = [%i, %i]" % (np.shape(h0)[0], np.shape(h0)[1]) , script_name = script_name)
  944. utils.print_verbose("h0.min max = [%0.15f, %0.15f]" % (np.amin(h0, axis = None), np.amax(h0, axis = None) ) , script_name = script_name)
  945. #print information about batch sampling
  946. utils.print_flush("batch_size: ( %i )" % (batch_size), script_name = script_name)
  947. utils.print_verbose("batch_size (number of subjects per batch): %i" % (batch_size), script_name = script_name)
  948. utils.print_flush("num_batch: ( %i )" % (num_batch), script_name = script_name)
  949. utils.print_verbose("num_batch (number of batches in full X): %i" % (num_batch), script_name = script_name)
  950. utils.print_flush("max_epoch: ( %i )" % (max_epoch), script_name = script_name)
  951. utils.print_flush("iteration_per_batch: ( %i )" % (iteration_per_batch), script_name = script_name)
  952. #initalize arrays to store values during the for loop of updates
  953. epoch_start = 0
  954. statistics_start_time = time.time()
  955. epoch_array = np.arange(start = 0, stop = max_epoch, step = 1)
  956. batch_loss_sum_per_epoch_array = np.zeros(shape = epoch_array.shape)
  957. full_loss_per_epoch_array = np.zeros(shape = epoch_array.shape)
  958. batch_loss_per_iteration_array = np.zeros(shape = (max_epoch, num_batch))
  959. full_loss_per_iteration_array = np.zeros(shape = (max_epoch, num_batch))
  960. sparsity_per_epoch_array = np.zeros(shape = epoch_array.shape)
  961. sparsity_per_iteration_array = np.zeros(shape = (max_epoch, num_batch))
  962. if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
  963. trXtX = np.power(np.linalg.norm(X, 'fro'), 2)
  964. obj = utils.get_objective_function(X = X, W = W, trXtX = trXtX, OPNMF = True)
  965. obj_old = obj
  966. objective_function = 0.0 #update to value if VERBOSE_FLAG is on
  967. statistics_elapsed_time = statistics_elapsed_time + (time.time() - statistics_start_time)
  968. del statistics_start_time
  969. #to facilitate |X-WH|_F^2 objective calculation, we want to store tr(X'X) in the memory and reuse it
  970. initialization_start_time = time.time()
  971. if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
  972. trXtX = np.power(np.linalg.norm(X, 'fro'), 2)
  973. obj = utils.get_objective_function(X = X, W = W, trXtX = trXtX, OPNMF = False) #because we cannot assume W to be orthogonal at the beginning, we cannot simplify obj as W'T=I assumption
  974. obj_old = obj
  975. initialization_elapsed_time = initialization_elapsed_time + (time.time() - initialization_start_time)
  976. del initialization_start_time
  977. ## load intermediate if it exists
  978. if os.path.isfile(initialization_path):
  979. saving_start_time = time.time()
  980. utils.print_flush(string_variable = "Loading initialization from file ( %s )." % (initialization_path), script_name = script_name)
  981. total_elapsed_time = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "total_elapsed_time", script_name = script_name)
  982. saving_elapsed_time = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "saving_elapsed_time", script_name = script_name)
  983. initialization_elapsed_time = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "initialization_elapsed_time", script_name = script_name)
  984. statistics_elapsed_time = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "statistics_elapsed_time", script_name = script_name)
  985. compression_elapsed_time = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "compression_elapsed_time", script_name = script_name)
  986. saving_elapsed_time = saving_elapsed_time + (time.time() - saving_start_time)
  987. del saving_start_time
  988. else:
  989. saving_start_time = time.time()
  990. mdict = {
  991. "w0": w0,
  992. "h0": h0,
  993. "num_batch": num_batch,
  994. "idx_sample": idx_sample,
  995. "total_elapsed_time": total_elapsed_time,
  996. "saving_elapsed_time": saving_elapsed_time,
  997. "initialization_elapsed_time": initialization_elapsed_time,
  998. "statistics_elapsed_time": statistics_elapsed_time,
  999. "compression_elapsed_time": compression_elapsed_time
  1000. }
  1001. if not COMPRESSION_PER_BATCH_FLAG:
  1002. mdict_temp = {"Q": Q}
  1003. mdict.update(mdict_temp)
  1004. del mdict_temp
  1005. utils.save_intermediate_output(output_path = initialization_path, mdict = mdict, OVERWRITE = False)
  1006. saving_elapsed_time = saving_elapsed_time + (time.time() - saving_start_time)
  1007. del saving_start_time
  1008. ## load from restart point
  1009. if os.path.isfile(restart_path):
  1010. utils.print_flush(string_variable = "Loading restart file from file ( %s )." % (restart_path), script_name = script_name)
  1011. epoch_start = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "epoch")
  1012. W = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "W")
  1013. W_old = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "W_old")
  1014. diffW = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "diffW")
  1015. total_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "total_elapsed_time")
  1016. initialization_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "initialization_elapsed_time")
  1017. saving_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "saving_elapsed_time")
  1018. statistics_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "statistics_elapsed_time")
  1019. update_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "update_elapsed_time")
  1020. sampling_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "sampling_elapsed_time")
  1021. compression_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "compression_elapsed_time")
  1022. utils.print_flush(string_variable = "Loading from restart: epoch ( %i ) and elapsed time ( %s seconds )." % (epoch_start, total_elapsed_time), script_name = script_name)
  1023. #start for loop
  1024. for epoch in np.arange(start = epoch_start, stop = epoch_array[-1], step = 1):
  1025. #restart happens at epoch level
  1026. if np.mod(epoch, restart_step) == 0:
  1027. saving_start_time = time.time()
  1028. # if epoch != 0:
  1029. # #to avoid filesystem io error (e.g. syncing across nodes or some other activity that may cause slight slowdown/lag and cause the code to fail to locate the file to move)
  1030. # time.sleep(0.5) #sleep for 0.5 seconds
  1031. # shutil.move(src = restart_path, dst = restart_path_old)
  1032. # time.sleep(0.5) #sleep for 0.5 seconds
  1033. utils.print_verbose(string_variable = "epoch %i / %i: saving intermediate save point for restarting." % (epoch, max_epoch), script_name = script_name)
  1034. mdict = {
  1035. "W": W,
  1036. "W_old": W_old,
  1037. "diffW": diffW,
  1038. "epoch": epoch,
  1039. "initialization_elapsed_time": initialization_elapsed_time,
  1040. "saving_elapsed_time": saving_elapsed_time,
  1041. "statistics_elapsed_time": statistics_elapsed_time,
  1042. "update_elapsed_time": update_elapsed_time,
  1043. "sampling_elapsed_time": sampling_elapsed_time,
  1044. "compression_elapsed_time": compression_elapsed_time,
  1045. "total_elapsed_time": total_elapsed_time + (time.time() - total_start_time)
  1046. }
  1047. utils.save_intermediate_output(output_path = restart_path, mdict = mdict, OVERWRITE = True)
  1048. # if not os.path.isfile(restart_path_old): #must be missing restart_path_old because epoch == 0 or it got deleted
  1049. # utils.save_intermediate_output(output_path = restart_path_old, mdict = mdict, OVERWRITE = False)
  1050. restart_path_epoch = os.path.join(outdir, "%s_restart_epoch%05d.mat" % (output_basename_prefix, epoch))
  1051. utils.print_verbose("Copying ( %s ) to ( %s ) as additional backup." % (restart_path, restart_path_epoch), script_name = script_name)
  1052. shutil.copy(src = restart_path, dst = restart_path_epoch)
  1053. restart_path_epoch_to_remove = os.path.join(outdir, "%s_restart_epoch%05d.mat" % (output_basename_prefix, epoch - 2 * restart_step))
  1054. if os.path.isfile(restart_path_epoch_to_remove):
  1055. utils.print_verbose("Removing very old restart file ( %s )." % (restart_path_epoch_to_remove), script_name = script_name)
  1056. os.unlink(restart_path_epoch_to_remove)
  1057. del restart_path_epoch_to_remove
  1058. del restart_path_epoch
  1059. del mdict
  1060. saving_elapsed_time = saving_elapsed_time + (time.time() - saving_start_time)
  1061. del saving_start_time
  1062. for batch_index in np.arange(num_batch): #e.g. 40 num batch, 1 iteration per sample, 50K max_epoch, which means there will be 25 samples per batch
  1063. ## Batch Sampling: select the batch of subjects from full X for this iteration
  1064. sampling_start_time = time.time()
  1065. utils.print_debug(string_variable = "epoch %i / %i: selecting %i-th batch with %i subjects." % (epoch, max_epoch, batch_index, batch_size), script_name = script_name)
  1066. idx_sample_selected = (idx_sample[:, int(batch_index), int(epoch)]).astype(int)
  1067. X_sampled = X[:, idx_sample_selected]
  1068. del idx_sample_selected
  1069. sampling_elapsed_time = sampling_elapsed_time + (time.time() - sampling_start_time)
  1070. del sampling_start_time
  1071. #compress X_sampled to X_sampled
  1072. if VERBOSE_FLAG:
  1073. print("X_sampled.shape = ", flush=True)
  1074. print(X_sampled.shape, flush=True)
  1075. #if COMPRESSION_PER_BATCH_FLAG is on, generate Q FROM batch X_sampled here to compress X_sampled
  1076. compression_start_time = time.time()
  1077. if COMPRESSION_PER_BATCH_FLAG:
  1078. if Q_X_compression_size < batch_size:
  1079. utils.print_flush("Q_X_compression_size ( %i ) < batch_size ( %i ). Resetting Q_X_compression_size to ( %i )." % (Q_X_compression_size, batch_size, batch_size), script_name = script_name)
  1080. Q_X_compression_size = batch_size
  1081. Q_full = get_Q_compression(X_hat = X_sampled, X_compression_size = Q_X_compression_size, power_iteration = power_iteration, l = l, axis = 1, VERBOSE_FLAG = VERBOSE_FLAG)
  1082. if VERBOSE_FLAG:
  1083. print("Q_full.shape = ", flush=True)
  1084. print(Q_full.shape, flush=True) #should be l by l?
  1085. #if Q_X_compression_size > batch_size, select first batch_size of Q_X_compression_size of Q to use for compression
  1086. if Q_X_compression_size < batch_size:
  1087. print("%s: ERROR - you generated Q compression matrix from %i subjects of X_hat but you are using mini-batches of size %i subjects, where %i > %i. Exiting." % (script_name, Q_X_compression_size, batch_size, Q_X_compression_size, batch_size), flush=True)
  1088. sys.exit(1)
  1089. elif Q_X_compression_size == batch_size:
  1090. Q = Q_full.copy()
  1091. else:
  1092. Q = Q_full[0:batch_size, :]
  1093. if VERBOSE_FLAG:
  1094. print("%s: Extracting Q (for actually compressing X_tilde mini-batch repeatedly) of shape [%i, %i] out of Q_full (from X_hat for generating Q) of shape [%i, %i]" % (script_name, Q.shape[0], Q.shape[1], Q_full.shape[0], Q_full.shape[1]), flush=True)
  1095. #Q_full is no longer needed now that we have final Q to use
  1096. del Q_full
  1097. if epoch == 0 and batch_index == 0:
  1098. utils.print_flush("COMPRESSION_PER_BATCH_FLAG: ( %r ) \t| Q.shape = [ %i, %i ]" % (COMPRESSION_PER_BATCH_FLAG, Q.shape[0], Q.shape[1]), script_name = script_name)
  1099. # X_sampled_compressed = np.matmul(Q.T, X_sampled) #if compressing m -> l
  1100. X_sampled_compressed = np.matmul(X_sampled, Q) #if compressing n -> l
  1101. if VERBOSE_FLAG:
  1102. print("X_sampled_compressed.shape = ",flush=True)
  1103. print(X_sampled_compressed.shape, flush=True)
  1104. compression_elapsed_time = compression_elapsed_time + ( time.time() - compression_start_time )
  1105. del compression_start_time
  1106. for iteration in np.arange(iteration_per_batch): #one iteration per batch
  1107. saving_start_time = time.time()
  1108. #save small file for plotting progress after the run finishes
  1109. if np.mod(epoch, save_step) == 0 and ( (batch_index == 0) and (iteration == 0) ):
  1110. current_elapsed_time = (time.time() - total_start_time) + (total_elapsed_time)
  1111. intermediate_debug_path = os.path.join(outdir, "%s_intermediate_epoch%05d.mat" % (output_basename_prefix, epoch))
  1112. utils.print_verbose(string_variable = "epoch: %i / %i \t| batch: %i / %i \t| iteration: %i / %i \t| diffW = %0.5E \t| elapsed time: %f seconds \t| saving save file ( %s )" % (epoch, max_epoch, batch_index, num_batch, iteration, iteration_per_batch, diffW, current_elapsed_time, intermediate_debug_path), script_name = script_name)
  1113. mdict = {
  1114. "W": W,
  1115. "elapsed_time": current_elapsed_time,
  1116. "epoch": epoch,
  1117. "iteration": iteration,
  1118. "batch_index": batch_index
  1119. }
  1120. hdf5storage.savemat(intermediate_debug_path, mdict = mdict)
  1121. del mdict, intermediate_debug_path
  1122. #print progress to console
  1123. if np.mod(epoch, print_step) == 0 and batch_index == 0 and iteration == 0:
  1124. current_elapsed_time = (time.time() - total_start_time) + (total_elapsed_time)
  1125. remaining_time_string = utils.get_remaining_time_string(iteration = epoch, max_iter = max_epoch, elapsed_time = (time.time() - total_start_time) + total_elapsed_time )
  1126. utils.print_flush(string_variable = "epoch: %i / %i \t| batch: %i / %i \t| iteration: %i / %i \t| diffW = %0.5E \t| elapsed time: %f seconds\t| remaining time (approx): %s" % (epoch, max_epoch, batch_index, num_batch, iteration, iteration_per_batch, diffW, current_elapsed_time, remaining_time_string), script_name = script_name)
  1127. del current_elapsed_time, remaining_time_string
  1128. saving_elapsed_time = saving_elapsed_time + (time.time() - saving_start_time)
  1129. del saving_start_time
  1130. utils.print_debug(string_variable = "Storing old W", script_name = script_name)
  1131. W_old = W
  1132. if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
  1133. utils.print_debug(string_variable = "Storing old objective function = %0.5E" % (obj), script_name = script_name)
  1134. obj_old = obj
  1135. #multiplicative update for W
  1136. update_start_time = time.time()
  1137. utils.print_debug(string_variable = "Update with %s multiplicative update rule on W" % (multiplicative_update_method), script_name = script_name)
  1138. if (multiplicative_update_method == "original") or (multiplicative_update_method == "normalize"):
  1139. #multiplicative update rule - modified version that does not require XX = X*X' ready and available
  1140. W = opnmf_update_rule.multiplicative_update(X = X_sampled_compressed, W = W, OPNMF = True, MEM = True, rho = 1, DEBUG = False, script_name = script_name)
  1141. elif (multiplicative_update_method == "decoupled"):
  1142. W = opnmf_update_rule.multiplicative_update_decoupled(W = W, X = X_sampled_compressed, script_name = script_name, SANITY_CHECK_FLAG = False)
  1143. elif (multiplicative_update_method == "adaptive") or (multiplicative_update_method == "adaptiveNormalize"):
  1144. W, obj, rho, REJECT_W = opnmf_update_rule.multiplicative_update_adaptive(W = W, X = X_sampled_compressed, trXtX = trXtX, obj = obj, rho = rho, rho0 = rho0, eta = eta, DEBUG = False, EPSILON = EPSILON)
  1145. elif (update_meth == "quadratic"):
  1146. W = opnmf_update_rule.multiplicative_update_quadratic(W = W, X = X_sampled_compressed, OPNMF = False, MEM = True, rho = rho0, EPSILON = EPSILON)
  1147. elif (update_meth == "quadraticOrthonormal"):
  1148. W = opnmf_update_rule.multiplicative_update_quadratic(W = W, X = X_sampled_compressed, OPNMF = True, MEM = True, rho = rho0, EPSILON = EPSILON)
  1149. else:
  1150. print("%s: multiplicative_update_method (%s) must be either original, normalize, decoupled, adaptive, adaptiveNormalize, quadratic, quadraticOrthonormal. Exiting." % (script_name, multiplicative_update_method), flush=True)
  1151. sys.exit(1)
  1152. #reset small values of W - to prevent slowdown of computations caused by very small values, call on reset_small_value function to reset element of W with values smaller than (default 1.0e-16) to (default 1.0e-16)
  1153. utils.print_debug(string_variable = "Resetting small values of W", script_name = script_name)
  1154. W = opnmf_update_rule.reset_small_value(W, min_reset_value = reset_small_value_threshold)
  1155. #normalize W following mulitplicative update to stabilize W; only if not using normalization mulitplicative update version, not for convergent (constant root) or adaptive update
  1156. if multiplicative_update_method == "normalize" or multiplicative_update_method == "adaptiveNormalize":
  1157. utils.print_debug(string_variable = "Normalizing W", script_name = script_name)
  1158. W = opnmf_update_rule.normalize_W(W = W)
  1159. if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
  1160. utils.print_debug(string_variable = "updating rho ( %0.5E ) for %s multiplicative update method." % (rho, multiplicative_update_method), flush=True)
  1161. bOK = False
  1162. while not bOK:
  1163. obj = utils.get_objective_function(X = X_sampled_compressed, W = W, trXtX = trXtX, OPNMF = False) #We cannot assume that the property (constraint) W'W = I is true at early iterations #not sure if I should use full X or X_sampled_compressed here for updating adaptive step size
  1164. if rho != rho0 and obj > obj_old:
  1165. bOK = False
  1166. rho = rho0
  1167. W = W_old
  1168. obj = obj_old
  1169. else:
  1170. bOK = True
  1171. rho = rho + eta
  1172. update_elapsed_time = update_elapsed_time + (time.time() - update_start_time)
  1173. del update_start_time
  1174. ## stopping criterion: diffW for normalization or convergent update using diffW
  1175. update_start_time = time.time()
  1176. diffW = np.linalg.norm(W_old - W, ord = 'fro') / np.linalg.norm(W_old, ord = 'fro')
  1177. utils.print_debug(string_variable = "diffW = %0.5E" % (diffW), script_name = script_name)
  1178. if multiplicative_update_method == "constant" or multiplicative_update_method == "normalize":
  1179. if diffW < tol:
  1180. utils.print_flush("Converged after %i / %i epochs, %i / %i batches, %i / %i iterations." % (epoch, max_epoch, batch_index, num_batch, iteration, iteration_per_batch), script_name = script_name)
  1181. update_elapsed_time = update_elapsed_time + (time.time() - update_start_time)
  1182. del update_start_time
  1183. ## Calculate H
  1184. update_start_time = time.time()
  1185. utils.print_flush(string_variable = "calculating final H", script_name = script_name)
  1186. H = np.matmul(np.transpose(W), X)
  1187. update_elapsed_time = update_elapsed_time + (time.time() - update_start_time)
  1188. del update_start_time
  1189. total_end_time = time.time()
  1190. total_end_time_string = datetime.datetime.fromtimestamp(total_end_time).strftime('%Y-%m-%d %H:%M:%S')
  1191. utils.print_flush(string_variable = "total start time\t: %s" % (total_start_time_string), script_name = script_name)
  1192. utils.print_flush(string_variable = "total end time\t: %s" % (total_end_time_string), script_name = script_name)
  1193. del total_start_time_string, total_end_time_string
  1194. total_elapsed_time = total_elapsed_time + (total_end_time - total_start_time)
  1195. #save final output
  1196. saving_start_time = time.time()
  1197. utils.print_flush(string_variable = "saving final output (%s)" % (output_path), script_name = script_name)
  1198. mdict = {
  1199. "X": X,
  1200. "W": W,
  1201. "H": H,
  1202. "Q": Q,
  1203. "batch_loss_per_iteration_array": batch_loss_per_iteration_array,
  1204. "full_loss_per_iteration_array": full_loss_per_iteration_array,
  1205. "sparsity_per_iteration_array": sparsity_per_iteration_array,
  1206. "batch_loss_sum_per_epoch_array": batch_loss_sum_per_epoch_array,
  1207. "full_loss_per_epoch_array": full_loss_per_epoch_array,
  1208. "sparsity_per_epoch_array": sparsity_per_epoch_array,
  1209. "total_elapsed_time": total_elapsed_time,
  1210. "initialization_elapsed_time": initialization_elapsed_time,
  1211. "saving_elapsed_time": saving_elapsed_time,
  1212. "statistics_elapsed_time": statistics_elapsed_time,
  1213. "update_elapsed_time": update_elapsed_time,
  1214. "sampling_elapsed_time": sampling_elapsed_time,
  1215. "compression_elapsed_time": compression_elapsed_time
  1216. }
  1217. utils.save_intermediate_output(output_path = output_path, mdict = mdict)
  1218. del mdict
  1219. saving_elapsed_time = saving_elapsed_time + (time.time() - saving_start_time)
  1220. del saving_start_time
  1221. utils.print_flush(string_variable = "Total optimization elapsed time: %f seconds" % (total_elapsed_time))
  1222. utils.print_flush(string_variable = "Saving/Restarting/IO: %f seconds" % (saving_elapsed_time))
  1223. utils.print_flush(string_variable = "Multiplicative Updates: %f seconds" % (update_elapsed_time))
  1224. utils.print_flush(string_variable = "Initialization: %f seconds" % (initialization_elapsed_time))
  1225. utils.print_flush(string_variable = "Statistics Calculations: %f seconds" % (statistics_elapsed_time))
  1226. utils.print_flush(string_variable = "Sampling: %f seconds" % (sampling_elapsed_time))
  1227. utils.print_flush(string_variable = "Compression: %f seconds" % (compression_elapsed_time))
  1228. sys.exit(0)

csopnmf.py at commit 7b40455, under Apache-2.0 · at the source

Overview

Authors: Abdalla Bani1, Sung Min Ha1, Thomas Earnest2, Braden Yang1, Pan Xiao1, John Lee1, Janine Bijsterbosch1, Aristeidis Sotiras1,2,3,4
  1. Mallinckrodt Institute of Radiology, Washington University in St. Louis, St. Louis, MO, United States
  2. Institute for Informatics, Data Science & Biostatistics (I, 2, DB), Washington University in St. Louis, St. Louis, MO, United States
  3. Department of Radiology, Perelman School of Medicine, University of Pennsylvania, Philadelphia, PA, United States
  4. Center for AI and Data Science for Integrated Diagnostics (AI2D), Perelman School of Medicine, University of Pennsylvania, Philadelphia, PA, United States
Institutions: Washington University in St. Louis (United States); University of Pennsylvania (United States)
Journal: Imaging neuroscience (Cambridge, Mass.), volume 4, article IMAG.a.1355
Dates: received 11 November 2025; accepted 9 August 2026; published online 8 September 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1162/imag.a.1355 · PMID 42719766 · PMCID PMC13556794 · OpenAlex W7203499477
Open access: diamond, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), human (organism)
Methods: Statistics, Machine learning, Preprocessing, Connectivity
Keywords: MRI, big data, NMF, data compression and stochastic optimization
MeSH: Big Data*, Brain*, Data Compression*, Neuroimaging*, Female, Gray Matter, Humans, Image Processing, Computer-Assisted, Machine Learning, Magnetic Resonance Imaging, Male, Principal Component Analysis, Stochastic Processes (* major topic)
Topic: Functional Brain Connectivity Studies (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: National Institute on Aging (NIH R01-AG067103)
Citations: not cited yet (Europe PMC); 70 references in the paper

Abstract

Large-scale neuroimaging datasets present remarkable opportunities for advancing our understanding of human brain structure and function. Data-driven pattern analysis methods such as Orthonormal Projective Non-negative Matrix Factorization (opNMF) are particularly well suited to uncover multivariate relationships within these data, offering greater interpretability and reproducibility than more conventional approaches such as principal component analysis (PCA) and independent component analysis (ICA). Despite its utility in clinical computational neuroscience, the application of opNMF in large cohort studies has been impeded by computational challenges and scalability limitations. In this work, we address these issues by introducing a stochastic optimization strategy that processes mini-batches of the data, substantially improving scalability. We further accelerate computation through random data compression and leverage repulsive point processes to diversify mini-batches, reducing redundancy and the variance of updates. We first evaluated our method on gray matter tissue density maps from 1,000 participants in the Open Access Series of Imaging Studies (OASIS). Compared with the original approach, it achieved similar approximation accuracy and factor interpretability while greatly reducing computational cost. To demonstrate practical utility, we then applied the framework to 10,000 participants from the UK Biobank, identifying 20 patterns of structural covariance (PSCs) and examined associations between visceral adipose tissue (VAT) and PSC loadings, finding significant relationships for 13 PSCs in females and 11 PSCs in males, most of which were negative. We further show how these patterns refine in higher-rank decompositions with 40 and 60 components. This enhanced opNMF framework opens new possibilities for large-scale neuroimaging analyses, facilitating deeper insights into brain structure in both health and disease.

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.

sotiraslab/csopNMF

License: Apache-2.0
State: the link answers, verified on 26 September 2026
Evidence: files inventoried
Commit: 7b404557b3c78761e9c67b96c2ed8da5e36ce44d, 14 September 2026
Languages: Python (7)
Size: 9 files, 7 scripts
Software Heritage: not archived
Found in: “Data and Code Availability”
Holds: README, license file
Not found: CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (7 files), SciPy (5 files), pandas (3 files), scikit-learn (2 files), h5py (1 file), NiBabel (1 file)
Availability: 1 check, the latest on 26 September 2026: the link answers
  • 26 September 2026: the link answers
9 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;
  • 7 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

No dataset and no data link were found in the paper.

Data and Code Availability

The code for this project is publicly available at https://github.com/sotiraslab/csopNMF.

The OASIS dataset used in this work is publicly available from the OASIS Brains repository. Access requires user registration and acceptance of the OASIS data use terms. All preprocessing steps and analysis parameters are described in the paper, and any additional details can be provided upon reasonable request. UK Biobank data are available to bona fide researchers by application to UK Biobank and approval of a research proposal. This study used data accessed under application number 47267. In accordance with UK Biobank’s material transfer and data access terms, we cannot publicly share the raw UKB data. Aggregated results and analysis are provided in the paper (and further details can be shared upon reasonable request) to enable reproducibility.

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, 27 September 2026: the first record

Recorded: type, language, journal, volume, pages, dates, 8 authors, 4 keywords, 13 MeSH terms, 1 funder, 67 references.

Cite

This paper

Bani, A., Ha, S. M., Earnest, T., Yang, B., Xiao, P., Lee, J., Bijsterbosch, J., & Sotiras, A. (2026). Taming dimensionality in big neuroimaging data: Efficient orthonormal projective NMF via stochastic learning and data compression. Imaging neuroscience (Cambridge, Mass.), 4, IMAG.a.1355. https://doi.org/10.1162/imag.a.1355

BibTeX

@article{bani2026taming,
author = {Bani, Abdalla and Ha, Sung Min and Earnest, Thomas and Yang, Braden and Xiao, Pan and Lee, John and Bijsterbosch, Janine and Sotiras, Aristeidis},
title = {{Taming dimensionality in big neuroimaging data: Efficient orthonormal projective NMF via stochastic learning and data compression}},
journal = {Imaging neuroscience (Cambridge, Mass.)},
year = {2026},
month = sep,
volume = {4},
pages = {IMAG.a.1355},
publisher = {MIT Press},
issn = {2837-6056},
doi = {10.1162/imag.a.1355},
url = {https://doi.org/10.1162/imag.a.1355},
pmid = {42719766},
pmcid = {PMC13556794}
}

RIS

TY - JOUR
AU - Bani, Abdalla
AU - Ha, Sung Min
AU - Earnest, Thomas
AU - Yang, Braden
AU - Xiao, Pan
AU - Lee, John
AU - Bijsterbosch, Janine
AU - Sotiras, Aristeidis
TI - Taming dimensionality in big neuroimaging data: Efficient orthonormal projective NMF via stochastic learning and data compression
T2 - Imaging neuroscience (Cambridge, Mass.)
J2 - Imaging Neurosci (Camb)
PY - 2026
DA - 2026/09/08
VL - 4
SP - IMAG.a.1355
SN - 2837-6056
PB - MIT Press
DO - 10.1162/imag.a.1355
UR - https://doi.org/10.1162/imag.a.1355
LA - en
ER -

CSL-JSON

{
"id": "10.1162/imag.a.1355",
"type": "article-journal",
"title": "Taming dimensionality in big neuroimaging data: Efficient orthonormal projective NMF via stochastic learning and data compression",
"container-title": "Imaging neuroscience (Cambridge, Mass.)",
"author": [
{
"family": "Bani",
"given": "Abdalla"
},
{
"family": "Ha",
"given": "Sung Min"
},
{
"family": "Earnest",
"given": "Thomas"
},
{
"family": "Yang",
"given": "Braden"
},
{
"family": "Xiao",
"given": "Pan"
},
{
"family": "Lee",
"given": "John"
},
{
"family": "Bijsterbosch",
"given": "Janine"
},
{
"family": "Sotiras",
"given": "Aristeidis"
}
],
"container-title-short": "Imaging Neurosci (Camb)",
"volume": "4",
"page": "IMAG.a.1355",
"DOI": "10.1162/imag.a.1355",
"PMID": "42719766",
"PMCID": "PMC13556794",
"ISSN": "2837-6056",
"publisher": "MIT Press",
"URL": "https://doi.org/10.1162/imag.a.1355",
"language": "en",
"issued": {
"date-parts": [
[
2026,
9,
8
]
]
}
}

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/s41467-026-72091-7 [code]
Coupled cross-sectional and longitudinal non-negative matrix factorization reveals dominant brain aging trajectories in 48,949 individuals.
Journal: Nature communications
In common: scikit-learn, pandas, NumPy, structural MRI / diffusion, 7 references
[2] doi:10.1038/s41467-026-73072-6 [code]
Mapping the spatiotemporal continuum of structural connectivity development across the human connectome in youth.
Journal: Nature communications
In common: NiBabel, pandas, NumPy, structural MRI / diffusion, 5 references
[3] doi:10.21203/rs.3.rs-9326213/v1 [code]
Multi-task fMRI outperforms resting-state fMRI for revealing task-invariant organization of the human brain
Journal: Research Square (preprint)
In common: h5py, NiBabel, scikit-learn, 3 other tools, 3 references
[4] doi:10.64898/2026.03.09.710558 [code]
Multi-task fMRI outperforms resting-state fMRI for revealing task-invariant organization of the human brain
Journal: bioRxiv (preprint)
In common: h5py, NiBabel, scikit-learn, 3 other tools, 3 references
[5] doi:10.1371/journal.pbio.3003856 [code]
Aging and metabolism contribute separately to brain-body health.
Journal: PLoS biology
In common: NiBabel, scikit-learn, pandas, 2 other tools, structural MRI / diffusion, 3 references
[6] doi:10.1038/s41467-026-75585-6 [code]
Brain network dynamics reflect psychiatric illness status and transdiagnostic symptom profiles across health and disease.
Journal: Nature communications
In common: h5py, scikit-learn, pandas, 2 other tools, 3 references
[7] doi:10.1073/pnas.2519586123 [code]
Personalized functional topography-based multisite brain age prediction modeling reveals divergent neurodevelopment in major depression.
Journal: Proceedings of the National Academy of Sciences of the United States of America
In common: scikit-learn, pandas, SciPy, 1 other tool, structural MRI / diffusion, 3 references
[8] doi:10.1038/s41467-026-71270-w [code]
Spatiotemporal dynamics of the human cortical functional hierarchy across the lifespan.
Journal: Nature communications
In common: h5py, NiBabel, scikit-learn, 3 other tools, 2 references
[9] doi:10.1038/s41597-026-07248-6 [code]
A large-scale fMRI dataset for vision-language semantic association.
Journal: Scientific data
In common: h5py, NiBabel, scikit-learn, 3 other tools, 2 references
[10] doi:10.7554/elife.107933 [code]
Modality-agnostic decoding of vision and language from fMRI.
Journal: eLife
In common: h5py, NiBabel, scikit-learn, 3 other tools, 2 references

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.