Taming dimensionality in big neuroimaging data: Efficient orthonormal projective NMF via stochastic learning and data compression.
The 2 matches
- [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] § 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
- ## Module Imports
- import sys #check python version, exit upon sanity check failures
- import os #path checks
- import fractions #for handling fractions in argparse
- import argparse #for taking in user input/parameter/switch
- import time #for measuring elapsed time
- import datetime #for default output path string formation
- import getpass #for getting username on compute jobs since os.getlogin() works in interactive jobs but not compute jobs
- import shutil #for renaming files
- import numpy as np #matrix multiplications and optimizations
- import pandas as pd #for reading csv
- import hdf5storage #for saving as hdf5 mat file (v7.3 on matlab); has much better compression than scipy savemat
- import utils #for loading and saving data
- import initialize_nmf #for W component initialization
- import opnmf_update_rule #for W component initialization
- import restart #for determining outdir
- ## Default Variables
- #default value for script name to be used for printing to console
- script_name=os.path.basename(__file__)
- script_dirname = os.path.dirname(__file__)
- username=str(getpass.getuser())
- #Version
- script_version = "20230911_160100"
- #sanity check: python version
- utils.check_python_version(major = 3, minor = 7) #ensure the user is using python3.7 or higher
- #FLAGS
- VERBOSE_FLAG = False
- DEBUG_FLAG = False
- WITH_REPLACEMENT_FLAG = False
- ## Default values to functions and argparser
- 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")
- feret_path = os.path.join("/scratch/sungminha/git/NMF_Testing/faces_1024_2409_min0_max1_double.mat")
- EPSILON = np.finfo(np.float32).eps
- ## cs-opNMF specific functions
- def left_compression_Q(X, k, k_ov, w):
- """
- X = data matrix
- k = target rank
- k_ov = oversampling, s.t. k + k_ov = l
- w = power iteration
- """
- m = X.shape[0]
- n = X.shape[1]
- omega = np.random.default_rng().normal(loc = 0, scale = 1, size = (n, l))
- B = np.matmul(X, omega)
- print("%s: left_compression_Q: B.shape = " % (script_name))
- print(B.shape, flush=True)
- for power_iteration in np.arange(w):
- B = np.matmul(X, np.matmul(X.T, B))
- Q, R = np.linalg.qr(B)
- return Q
- def right_compression_Q(X, k, k_ov, w):
- """
- X = data matrix
- k = target rank
- k_ov = oversampling, s.t. k + k_ov = l
- w = power iteration
- """
- m = X.shape[0]
- n = X.shape[1]
- omega = np.random.default_rng().normal(loc = 0, scale = 1, size = (l, m))
- B = np.matmul(omega, X)
- print("%s: right_compression_Q: B.shape = " % (script_name))
- print(B.shape, flush=True)
- for power_iteration in np.arange(w):
- B = np.matmul(np.matmul(B, X.T), X)
- Q, R = np.linalg.qr(B.T)
- return Q.T
- def get_Q_compression(X_hat, X_compression_size = 100, power_iteration = 4, l = 50, axis = 1,VERBOSE_FLAG = False ):
- """
- axis = 0: Get Q such that it compresses feature space of X: X[m,n] to X[l,n]
- axis = 1: Get Q such that it compresses subject space of X: X[m,n] to X[m,l]
- where we assume:
- l = np.amin([n, np.amax([compression_level, target_rank + 10])]) #r + r_ov
- X_compression_size: How many subjects of X to use for calculating Q
- e.g.) X[m,n], X_compression_size = n: then use full X
- 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
- recommended to use full X by providing X_compression_size = n if X matrix is not too large for computational resources
- l: output dimension after compression for the axis of interest
- axis = 0: X[m,n] -> X_compressed[l,n]
- axis = 1: X[m,n] -> X_compressed[m,l]
- axis = 0:
- Q = np.random.rand(m, l) shape of Q must be m by l
- omega = np.random.standard_normal(size = (n, l))
- axis = 1:
- Q = np.random.rand(l, n)
- omega = np.random.standard_normal(size = (l, m))
- """
- #sanity check: axis parameter must be 0 or 1 assuming 2D input data matrix
- if axis != 0 and axis != 1:
- print("%s: get_Q_compression: axis parameter must be 0 or 1, not %0.5E. Exiting." % (script_name, axis), flush=True)
- sys.exit(1)
- #sanity check: X must be 2D
- X_hat_shape = np.shape(X_hat)
- if np.shape(X_hat_shape)[0] != 2:
- 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)
- sys.exit(1)
- m = np.shape(X_hat)[0]
- n = np.shape(X_hat)[1]
- if VERBOSE_FLAG:
- print("%s: get_Q_compression: data matrix provided to calculate X_hat.shape = [%i, %i]." % (script_name, m, n), flush=True)
- #generate omega matrix accordingly to axis parameter
- if axis == 0:
- omega = np.random.default_rng().normal(loc = 0, scale = 1, size = (X_compression_size, l))
- elif axis == 1:
- omega = np.random.default_rng().normal(loc = 0, scale = 1, size = (l, m))
- else:
- print("%s: get_Q_compression: unknown value ( %0.5E ) for axis parameter. Exiting." % (script_name, axis),flush=True)
- sys.exit(1)
- if VERBOSE_FLAG:
- print("%s: omega.shape = " % (script_name), flush=True)
- print(omega.shape, flush=True)
- #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
- if VERBOSE_FLAG:
- print("%s: get_Q_compression: X_compression_size = %i" % (script_name, X_compression_size), flush=True)
- if n < X_compression_size:
- 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)
- sys.exit(1)
- elif n == X_compression_size:
- X_compression_sample = X_hat
- else: #n > X_compression_size - need to sample mini-batch from X_hat
- idx_list = np.arange(start = 0, stop = n, step = 1) #list of indices to sample from
- X_compression_idx = np.random.choice(idx_list, size = X_compression_size, replace=False) #p=None default -> uniform random sampling
- X_compression_sample = X_hat[:,X_compression_idx]
- del X_compression_idx
- if VERBOSE_FLAG:
- print("%s: X_compression_sample.shape = " % (script_name), flush=True)
- print(X_compression_sample.shape, flush=True)
- if axis == 0:
- B = np.matmul(X_compression_sample, omega)
- #power iterations
- for j in np.arange(power_iteration): #same as (XX')^(power iteration)W
- B = np.matmul(X_compression_sample, np.matmul(X_compression_sample.T, B))
- Q, R = np.linalg.qr(B)
- elif axis == 1:
- B = np.matmul(omega, X_compression_sample) #l, n_batch
- #power iterations
- for j in np.arange(power_iteration): #same as H(X'X)^(power iteration)
- B = np.matmul(np.matmul(B, X_compression_sample.T), X_compression_sample )
- #get QR decomposition
- if axis == 0:
- Q, R = np.linalg.qr(B)
- elif axis == 1:
- Q, R = np.linalg.qr(B.T)
- else:
- print("%s: get_Q_compression: unknown value ( %0.5E ) for axis parameter. Exiting." % (script_name, axis),flush=True)
- sys.exit(1)
- if VERBOSE_FLAG:
- print("%s: get_Q_compression: Q.shape after QR decomposition = " % (script_name), flush=True)
- print(Q.shape, flush=True)
- print("%s: get_Q_compression: R.shape after QR decomposition = " % (script_name), flush=True)
- print(R.shape, flush=True)
- return Q
- def get_reconstruction_error(X, Q, VERBOSE_FLAG = False ):
- """
- reconstruction error in terms of compression; how well does compressed X reconstruct back to original full X
- """
- if VERBOSE_FLAG:
- print("%s: get_reconstruction_error: X.shape = " % (script_name), flush=True)
- print(X.shape, flush=True)
- print("%s: get_reconstruction_error: Q.shape = " % (script_name), flush=True)
- print(Q.shape, flush=True)
- return np.power(np.linalg.norm(X - np.matmul(Q, np.matmul(Q.T, X))), 2)
- ## s-opNMF specific functions
- def sample_uniform(X, num_batch, batch_size, max_epoch, axis = 1, WITH_REPLACEMENT = False):
- """
- inputs
- X: input matrix of size [m,n]
- axis (default=1): which direction to sample from; default is n in X in [m,n] where n is assumed to be subjects
- """
- #assign indices from 0 to n (where n is number of subjects, or number of columns of X)
- idx_list = np.arange(start = 0, stop = np.shape(X)[axis], step = 1)
- 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
- for epoch in np.arange(max_epoch):
- if VERBOSE_FLAG:
- if np.mod(epoch, 100) == 0: #don't print out all epochs, will print out too many too often
- print("%s: sample_uniform: sampling %i-th epoch indices." % (script_name, epoch), flush=True)
- idx_sample_full[:,:,epoch] = np.random.choice(idx_list, size = (batch_size, num_batch), replace = WITH_REPLACEMENT)
- return idx_sample_full
- def sample_dpp(evalue, evector, k):
- """
- sample a set Y from a dpp. evalue, evector are a decomposed kernel, and k is (optionally) the size of the set to return
- :param evalue: eigenvalue
- :param evector: normalized eigenvector
- :param k: number of cluster
- :return:
- """
- if k == None:
- # choose eigenvectors randomly
- evalue = np.divide(evalue, (1 + evalue))
- v = np.where(np.random.random(evalue.shape[0]) <= evalue)[0]
- #evector = np.where(np.random.random(evalue.shape[0]) <= evalue)[0]
- else:
- v = sample_k(evalue, k) ## v here is a 1d array with size: k
- k = v.shape[0]
- v = v.astype(int)
- v = [i - 1 for i in v.tolist()] ## due to the index difference between matlab & python, here, the element of v is for matlab
- V = evector[:, v]
- ## iterate
- y = np.zeros(k)
- for i in range(k, 0, -1):
- ## compute probabilities for each item
- P = np.sum(np.square(V), axis=1)
- P = P / np.sum(P)
- # choose a new item to include
- y[i-1] = np.where(np.random.rand(1) < np.cumsum(P))[0][0]
- y = y.astype(int)
- # choose a vector to eliminate
- j = np.where(V[y[i-1], :])[0][0]
- Vj = V[:, j]
- V = np.delete(V, j, 1)
- ## Update V
- if V.size == 0:
- pass
- else:
- V = np.subtract(V, np.multiply(Vj, (V[y[i-1], :] / Vj[y[i-1]])[:, np.newaxis]).transpose()) ## watch out the dimension here
- ## orthogonalize
- for m in range(i - 1):
- for n in range(m):
- V[:, m] = np.subtract(V[:, m], np.matmul(V[:, m].transpose(), V[:, n]) * V[:, n])
- V[:, m] = V[:, m] / np.linalg.norm(V[:, m])
- y = np.sort(y)
- return y
- def sample_k(lambda_value, k):
- """
- Pick k lambdas according to p(S) \propto prod(lambda \in S)
- :param lambda_value: the corresponding eigenvalues
- :param k: the number of clusters
- :return:
- """
- ## compute elementary symmetric polynomials
- E = elem_sym_poly(lambda_value, k)
- ## ietrate over the lambda value
- num = lambda_value.shape[0]
- remaining = k
- S = np.zeros(k)
- while remaining > 0:
- #compute marginal of num given that we choose remaining values from 0:num-1
- if num == remaining:
- marg = 1
- else:
- marg = lambda_value[num-1] * E[remaining-1, num-1] / E[remaining, num]
- # sample marginal
- if np.random.rand(1) < marg:
- S[remaining-1] = num
- remaining = remaining - 1
- num = num - 1
- return S
- def elem_sym_poly(lambda_value, k):
- """
- given a vector of lambdas and a maximum size k, determine the value of
- the elementary symmetric polynomials:
- E(l+1,n+1) = sum_{J \subseteq 1..n,|J| = l} prod_{i \in J} lambda(i)
- :param lambda_value: the corresponding eigenvalues
- :param k: number of clusters
- :return:
- """
- N = lambda_value.shape[0]
- E = np.zeros((k + 1, N + 1))
- E[0, :] = 1
- for i in range(1, k+1):
- for j in range(1, N+1):
- E[i, j] = E[i, j - 1] + lambda_value[j-1] * E[i - 1, j - 1]
- return E
- def sample_dpp_linear(X, num_batch, batch_size, max_epoch, data, axis = 1, print_step = -1, WITH_REPLACEMENT = False, sampling_intermediate_path = "NA"):
- import scipy.stats
- import scipy.linalg
- #assume data is csv of age and sex, is in order of generating the input matrix X
- age_sex = data[["age", "sex"]]
- age_sex['sex'].replace(['F','M'],[-1,1],inplace=True)
- age_sex=age_sex.apply(scipy.stats.zscore) #z-score meta-data
- #extract mete-data vectors to create kernel
- age = age_sex['age']
- age = np.reshape([age],(len(age),1))
- age_transpose = np.transpose(age)
- sex = age_sex['sex']
- sex = np.reshape([sex],(len(sex),1))
- sex_transpose = np.transpose(sex)
- kern_original = np.dot(X.transpose(), X) #Linear covariance matrix based on imaging features
- idx_array_original = np.arange(start = 0, stop = X.shape[axis], step = 1)
- # 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
- idx_sample_full = -1 * np.ones(shape = (batch_size, num_batch, max_epoch)) #was originally in wrong shape
- finished_epoch = 0
- #kernel does not change each epoch; no need to repeat this each epoch; run once and reuse evalue and evector
- evalue_original, evector_original = scipy.linalg.eigh(kern_original)
- if os.path.isfile(sampling_intermediate_path):
- print("sample_dpp_linear: loading sampling indices from sampling_intermediate_path ( %s )." % (sampling_intermediate_path), flush=True)
- idx_sample_full = hdf5storage.loadmat(file_name = sampling_intermediate_path, variable_names = ["idx_sample_full"])["idx_sample_full"]
- finished_epoch = hdf5storage.loadmat(file_name = sampling_intermediate_path, variable_names = ["finished_epoch"])["finished_epoch"]
- 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)
- if print_step < 0:
- print_step = max_epoch #reset to print none to console
- for epoch in np.arange(start = finished_epoch, stop = max_epoch, step = 1):
- if np.mod(epoch, print_step) == 0:
- 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)
- #save intermediate
- sampling_intermediate_dir = os.path.dirname(sampling_intermediate_path)
- if not os.path.isdir(sampling_intermediate_dir):
- print("sample_dpp_linear: sampling_intermediate_dir ( %s ) does not exist. Not saving intermediates." % (sampling_intermediate_dir), flush=True)
- else:
- if np.mod(epoch, 100) == 0:
- mdict = {
- "idx_sample_full": idx_sample_full,
- "finished_epoch": epoch
- }
- print("sample_dpp_linear: Saving intermediate sampling to ( %s )" % (sampling_intermediate_path), flush=True)
- hdf5storage.savemat(file_name = sampling_intermediate_path, mdict=mdict)
- specific_sampling_intermediate_basename = os.path.splitext(os.path.basename(sampling_intermediate_path))[0]
- specific_sampling_intermediate_path = os.path.join(sampling_intermediate_dir, "%s_epoch%05d.mat" % (specific_sampling_intermediate_basename, epoch))
- print("sample_dpp_linear: Saving intermediate sampling to ( %s )" % (specific_sampling_intermediate_path), flush=True)
- hdf5storage.savemat(file_name = specific_sampling_intermediate_path, mdict=mdict)
- # idx_sample = -1 * np.ones((num_batch,batch_size))
- idx_sample = -1 * np.ones(shape = (batch_size, num_batch)) #was originally in wrong shape
- idx_array = np.copy(idx_array_original)
- kern = np.copy(kern_original)
- evector = np.copy(evector_original)
- evalue = np.copy(evalue_original)
- for sample in range(num_batch):
- if VERBOSE_FLAG:
- 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)
- idx = sample_dpp( np.real(evalue),np.real(evector), batch_size)
- idx_sample[:, sample] = idx_array[idx]
- if not WITH_REPLACEMENT:
- kern = np.delete(kern, idx, 0)
- kern = np.delete(kern, idx, 1)
- idx_array = np.delete(idx_array, idx, 0)
- # need to also update evalue and evector
- evalue = np.delete(evalue, idx, 0)
- evector = np.delete(evector, idx, 0)
- evector = np.delete(evector, idx, 1)
- idx_sample_full[:,:,epoch] = idx_sample
- del idx_sample
- #sanity check
- if np.sum(idx_sample_full < 0):
- 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)
- return idx_sample_full
- def sample_dpp_gaussian(X, num_batch, batch_size, max_epoch, data, sigma = 0.1, axis = 1, print_step = -1, WITH_REPLACEMENT = False):
- import scipy.stats
- age_sex = data[["age", "sex"]]
- age_sex['sex'].replace(['F','M'],[-1,1],inplace=True)
- age_sex=age_sex.apply(scipy.stats.zscore) #z-score meta-data
- #extract mete-data vectors to create kernel
- age = age_sex['age']
- age = np.reshape([age],(len(age),1))
- age_transpose = np.transpose(age)
- sex = age_sex['sex']
- sex = np.reshape([sex],(len(sex),1))
- sex_transpose = np.transpose(sex)
- #kernel
- #to-do: make into a function that takes some meta data (optional and required0 and generates a kernel
- kern_original = np.exp(-((age-age_transpose)**2+(sex-sex_transpose)**2)/sigma**2) #gaussian
- idx_array_original = np.arange(start = 0, stop = X.shape[axis], step = 1)
- # 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
- idx_sample_full = -1 * np.ones(shape = (batch_size, num_batch, max_epoch)) #was originally in wrong shape
- #kernel does not change each epoch; no need to repeat this each epoch; run once and reuse evalue and evector
- evalue_original, evector_original = scipy.linalg.eigh(kern_original)
- if print_step < 0:
- print_step = max_epoch #reset to print none to console
- for epoch in np.arange(max_epoch):
- if np.mod(epoch, print_step) == 0:
- 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)
- idx_array = np.copy(idx_array_original)
- kern = np.copy(kern_original)
- evector = np.copy(evector_original)
- evalue = np.copy(evalue_original)
- # idx_sample = -1 * np.ones((num_batch,batch_size))
- idx_sample = -1 * np.ones(shape = (batch_size, num_batch)) #was originally in wrong shape
- for sample in range(num_batch):
- if VERBOSE_FLAG:
- 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)
- idx = sample_dpp(np.real(evalue),np.real(evector), batch_size)
- # idx_sample[sample,:] = idx_array[idx]
- idx_sample[:,sample] = idx_array[idx]
- if not WITH_REPLACEMENT:
- kern = np.delete(kern, idx, 0)
- kern = np.delete(kern, idx, 1)
- idx_array = np.delete(idx_array, idx, 0)
- # need to also update evalue and evector
- evalue = np.delete(evalue, idx, 0)
- evector = np.delete(evector, idx, 0)
- evector = np.delete(evector, idx, 1)
- idx_sample_full[:,:,epoch] = idx_sample
- del idx_sample
- #sanity check
- if np.sum(idx_sample_full < 0):
- 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)
- return idx_sample_full
- def calculate_batch_loss_sum(X, idx_sample, W, SQUARED = True):
- """
- 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
- sum_{i=0}^{last batch}(|X_batch - W_final (W_final * X_batch)|_F^2)
- """
- batch_loss_sum = 0.0
- num_batch = idx_sample.shape[1]
- for batch_index in np.arange(num_batch):
- idx_sample_selected = (idx_sample[:, int(batch_index)]).astype(int)
- X_sampled = X[:, idx_sample_selected]
- del idx_sample_selected
- #for the sake of consistency, let us call the calculate_error function from utils module rather than calling frobenius norm function from scratch
- batch_loss_sum = batch_loss_sum + utils.calculate_error(X = X_sampled, W = W, H = np.matmul(W.T, X_sampled), SQUARED = SQUARED)
- del X_sampled
- return batch_loss_sum
- ## argparser
- parser=argparse.ArgumentParser(
- 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.",
- epilog = "Written by Sung Min Ha ([email hidden])",
- add_help=True,
- )
- parser.add_argument("-i", "--inputFile",
- type = str,
- dest = "input_path",
- required = False,
- default = feret_path,
- 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."
- )
- parser.add_argument("-d", "--demographic_data_path",
- type = str,
- dest = "demographic_data_path",
- required = False,
- default = csv_path,
- 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)."
- )
- parser.add_argument("-k", "--targetRank",
- type = int,
- dest = "target_rank",
- required = False,
- default = 40,
- help = "What is the rank of the NMF you intend to run? This is the number of components that will be generated."
- )
- parser.add_argument("-m", "--maxEpoch",
- type = int,
- dest = "max_epoch",
- required = False,
- default = 5.0e4,
- 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."
- )
- parser.add_argument("-t", "--tol",
- type = float,
- dest = "tol",
- required = False,
- default = 0.0e0,
- 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."
- )
- parser.add_argument("-o", "--outputParentDir",
- type = str,
- dest = "output_parent_dir",
- required = False,
- default = os.path.join("/scratch/%s" % (username), "output_directory"),
- help = "Path to output mat file that contain the outputs."
- )
- parser.add_argument("-0", "--initMeth",
- type = str,
- dest = "init_method",
- required = False,
- default = "nndsvd",
- help = "Method for initializing w0 for component. (random, nndsvd, nndsvda, nndsvdar)"
- )
- parser.add_argument("-u", "--updateMeth",
- type = str,
- dest = "update_meth",
- required = False,
- default = "mem",
- help = "Method for update W. (mem, original)"
- )
- parser.add_argument("--multiplicativeUpdateMeth",
- type = str,
- dest = "multiplicative_update_method",
- required = False,
- default = "normalize",
- help = "Method for multiplicative update of W. (normalize, constant, adaptive, adaptiveNormalize, quadratic, quadraticOrthonormal)"
- )
- parser.add_argument("-s", "--samplingMeth",
- type = str,
- dest = "sampling_method",
- required = False,
- default = "uniform",
- help = "Method for sampling of batch. (uniform, dppgaussian, dpplinear)"
- )
- parser.add_argument("-z", "--sampleSize",
- type = int,
- dest = "batch_size",
- required = False,
- default = 100,
- help = "Number of subjects per batch"
- )
- parser.add_argument("--printStep",
- type = int,
- dest = "print_step",
- required = False,
- default = 1.0e1,
- help = "Print progress every this many epochs."
- )
- parser.add_argument("--saveStep",
- type = int,
- dest = "save_step",
- required = False,
- default = 1.0e3,
- 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."
- )
- parser.add_argument("--restartStep",
- type = int,
- dest = "restart_step",
- required = False,
- default = 1.0e0,
- 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."
- )
- parser.add_argument("--rho",
- type = float,
- dest = "rho",
- required = False,
- default = 0.25,
- help = "For constant power or adaptive multiplicative update."
- )
- parser.add_argument("--eta",
- type = float,
- dest = "eta",
- required = False,
- default = 0.1,
- help = "For adaptive multiplicative update."
- )
- parser.add_argument("--sigma",
- type = float,
- dest = "sigma",
- required = False,
- default = 0.1,
- help = "For dpp gaussian kernel generation."
- )
- parser.add_argument("-V", "--verbose",
- action = 'store_true',
- dest = "VERBOSE_FLAG",
- help = "Extra printouts for debugging."
- )
- parser.add_argument("-D", "--debug",
- action = 'store_true',
- dest = "DEBUG_FLAG",
- help = "EXTRA EXTRA printouts for debugging."
- )
- parser.add_argument("--withReplacement",
- action = 'store_true',
- dest = "WITH_REPLACEMENT_FLAG",
- help = "If set to true by calling this flag on, DPP sampling will happen with replacement instead of removing already selected sample"
- )
- parser.add_argument("--QCompressionBatchSize",
- type = int,
- dest = "Q_X_compression_size",
- required = False,
- default = 1.0e3,
- 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)."
- )
- parser.add_argument("--oversampling",
- type = int,
- dest = "oversampling",
- required = False,
- default = 10,
- 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"
- )
- parser.add_argument("--powerIteration",
- type = int,
- dest = "power_iteration",
- required = False,
- default = 4,
- help = "What is the power iteration for compression matrix Q generation? Defaults to 4."
- )
- parser.add_argument("--compressionPerBatch",
- action = 'store_true',
- dest = "COMPRESSION_PER_BATCH_FLAG",
- 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."
- )
- parser.add_argument("--notOrthonormal",
- action = 'store_false',
- dest = "ORTHONORMAL_FLAG",
- help = "Use the orthonormal projective update instead of projective update"
- )
- parser.add_argument("--speedOptimized",
- action = 'store_false',
- dest = "MEM_FLAG",
- help = "Do you want to use speed optimized but memory inefficient multiplicative update where you store XX^T in the memory and reuse it."
- )
- parser.add_argument("--noZeroRemoval",
- action = 'store_false',
- dest = "ZERO_REMOVAL_FLAG",
- help = "turn off the flag to remove all rows in the X input matrix where it is zero across the entire row."
- )
- parser.add_argument("--noSmallValueReset",
- action = 'store_false',
- dest = "SMALL_VALUE_RESET_FLAG",
- 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)?"
- )
- parser.add_argument("--noSmallValueResetInit",
- action = 'store_false',
- dest = "SMALL_VALUE_RESET_INIT_FLAG",
- 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?"
- )
- parser.add_argument("--minResetValue",
- type = str,
- dest = "min_reset_value",
- default = "1.0e-16",
- 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."
- )
- parser.add_argument("--epsilon",
- type = str,
- dest = "EPSILON",
- default = str(np.finfo(np.float64).eps),
- 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."
- )
- parser.add_argument("--iterationPerBatch",
- type = int,
- dest = "iteration_per_batch",
- default = int(1),
- help = "How many iterations of multiplicative updates do you want to perform per batch per epoch? default is 1."
- )
- #parse argparser
- args = parser.parse_args()
- input_path = args.input_path
- demographic_data_path = args.demographic_data_path
- target_rank = args.target_rank
- tol = args.tol #tolerance
- output_parent_dir = args.output_parent_dir
- init_method = args.init_method
- sampling_method = args.sampling_method
- multiplicative_update_method = args.multiplicative_update_method
- batch_size = args.batch_size
- update_meth = args.update_meth
- print_step = args.print_step
- save_step = args.save_step
- restart_step = args.restart_step
- VERBOSE_FLAG = args.VERBOSE_FLAG
- DEBUG_FLAG = args.DEBUG_FLAG #also save all intermediates at save point
- WITH_REPLACEMENT_FLAG = args.WITH_REPLACEMENT_FLAG
- # iter_power = args.iter_power
- max_epoch = int(args.max_epoch) #same as number of iterations for opNMF, e.g.) 50K
- rho = args.rho #1/4 by default
- eta = args.eta #0.1 by default
- sigma = args.sigma #0.1 by default
- SQUARED_ERROR_FLAG = True
- iteration_per_batch = int(args.iteration_per_batch)
- reset_small_value_threshold = 1.0e-16
- Q_X_compression_size = args.Q_X_compression_size
- oversampling = args.oversampling
- power_iteration = args.power_iteration
- COMPRESSION_PER_BATCH_FLAG = args.COMPRESSION_PER_BATCH_FLAG
- #update related parameters
- MEM_FLAG = args.MEM_FLAG
- ORTHONORMAL_FLAG = args.ORTHONORMAL_FLAG
- #additional resetting/small value/zero handling parameters
- SMALL_VALUE_RESET_FLAG = args.SMALL_VALUE_RESET_FLAG
- SMALL_VALUE_RESET_INIT_FLAG = args.SMALL_VALUE_RESET_INIT_FLAG
- min_reset_value = float(fractions.Fraction(str(args.min_reset_value)))
- if SMALL_VALUE_RESET_INIT_FLAG and not SMALL_VALUE_RESET_FLAG:
- 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
- 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 )
- EPSILON = float(fractions.Fraction(str(args.EPSILON)))
- ZERO_REMOVAL_FLAG = args.ZERO_REMOVAL_FLAG
- if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
- #store original rho
- rho0 = rho
- ORTHONORMAL_FLAG = True
- #sanity check: check whether correct configuration is setup for speed vs mem
- MEM_FLAG = False
- if (update_meth == "mem"):
- pass
- MEM_FLAG = True
- elif (update_meth == "original"):
- pass
- else:
- utils.print_flush("ERROR - Unknown update_meth (%s). Accepted values are mem and original. Exiting." % (update_meth), script_name = script_name)
- sys.exit(1)
- #sanity check: turn verbose flag on if debug flag is on
- if DEBUG_FLAG:
- VERBOSE_FLAG = True
- #print out final variables after argparsing
- print("\n\n", flush=True)
- utils.print_flush(string_variable = "-----Variables-----\n\n", script_name = script_name )
- utils.print_flush(string_variable = "input_path: ( %s )" % (input_path), script_name = script_name)
- utils.exit_if_not_exist_file(input_path)
- utils.print_flush(string_variable = "target_rank: ( %i )" % (target_rank), script_name = script_name)
- utils.print_flush(string_variable = "tol: ( %0.5E )" % (tol), script_name = script_name)
- utils.print_flush(string_variable = "output_parent_dir: ( %s )" % (output_parent_dir), script_name = script_name)
- utils.print_flush(string_variable = "init_method: ( %s )" % (init_method), script_name = script_name)
- utils.print_flush(string_variable = "sampling_method: ( %s )" % (sampling_method), script_name = script_name)
- utils.print_flush(string_variable = "multiplicative_update_method: ( %s )" % (multiplicative_update_method), script_name = script_name)
- utils.print_flush(string_variable = "ORTHONORMAL_FLAG: ( %r )" % (ORTHONORMAL_FLAG), script_name = script_name)
- utils.print_flush(string_variable = "MEM_FLAG: ( %r )" % (MEM_FLAG), script_name = script_name)
- utils.print_flush(string_variable = "update_meth: ( %s )" % (update_meth), script_name = script_name)
- utils.print_flush(string_variable = "print_step: ( %i )" % (print_step), script_name = script_name)
- utils.print_flush(string_variable = "restart_step: ( %i )" % (restart_step), script_name = script_name)
- utils.print_flush(string_variable = "save_step: ( %i )" % (save_step), script_name = script_name)
- utils.print_flush(string_variable = "max_epoch: ( %i )" % (max_epoch), script_name = script_name)
- utils.print_flush(string_variable = "rho: ( %0.5E )" % (rho), script_name = script_name)
- if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
- utils.print_flush(string_variable = "rho0: ( %0.5E )" % (rho0), script_name = script_name)
- if sampling_method == "dppgaussian":
- utils.print_flush(string_variable = "sigma: ( %0.5E )" % (sigma), script_name = script_name)
- utils.print_flush(string_variable = "eta: ( %0.5E )" % (eta), script_name = script_name)
- #compression related values
- utils.print_flush(string_variable = "batch_size: ( %i )" % (batch_size), script_name = script_name)
- utils.print_flush(string_variable = "WITH_REPLACEMENT_FLAG: ( %r )" % (WITH_REPLACEMENT_FLAG), script_name = script_name)
- utils.print_flush(string_variable = "Q_X_compression_size: ( %i )" % (Q_X_compression_size), script_name = script_name)
- utils.print_flush(string_variable = "oversampling: ( %i )" % (oversampling), script_name = script_name)
- utils.print_flush(string_variable = "power_iteration: ( %i )" % (power_iteration), script_name = script_name)
- utils.print_flush(string_variable = "COMPRESSION_PER_BATCH_FLAG: ( %r )" % (COMPRESSION_PER_BATCH_FLAG), script_name = script_name)
- utils.print_flush(string_variable = "EPSILON: ( %0.2E )" % (EPSILON), script_name = script_name)
- utils.print_verbose(string_variable = "EPSILON precision: %i" % (utils.get_numpy_precision(EPSILON)), script_name = script_name, VERBOSE_FLAG = VERBOSE_FLAG)
- utils.print_flush(string_variable = "SMALL_VALUE_RESET_INIT_FLAG: ( %r )" %
- (SMALL_VALUE_RESET_INIT_FLAG), script_name = script_name)
- 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)
- utils.print_flush(string_variable = "SMALL_VALUE_RESET_FLAG: ( %r )" % (SMALL_VALUE_RESET_FLAG), script_name = script_name)
- 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)
- utils.print_flush(string_variable = "min_reset_value: ( %0.2E )" % (min_reset_value), script_name = script_name)
- utils.print_flush(string_variable = "ZERO_REMOVAL_FLAG: ( %r )" % (ZERO_REMOVAL_FLAG), script_name = script_name)
- utils.print_flush(string_variable = "VERBOSE_FLAG: ( %r )" % (VERBOSE_FLAG), script_name = script_name)
- utils.print_flush(string_variable = "DEBUG_FLAG: ( %r )" % (DEBUG_FLAG), script_name = script_name)
- if ORTHONORMAL_FLAG:
- nmf_method = "opNMF"
- else:
- nmf_method = "pNMF"
- nmf_method = "cs-%s" % (nmf_method) #cs-opNMF or cs-pNMF
- ## initialize elapse time counts
- total_elapsed_time = 0.0 #all time elapsed from this point on
- initialization_elapsed_time = 0.0 #for initializations
- saving_elapsed_time = 0.0 #for intermediate savings and loading from restart
- statistics_elapsed_time = 0.0 #for calculating statistics (error/sparsity on the fly)
- update_elapsed_time = 0.0 #for multiplicative update
- sampling_elapsed_time = 0.0 #for sampling batches and selecting those batches as indices
- compression_elapsed_time = 0.0 #for compression (calculation and applying of Q matrices)
- total_start_time = time.time()
- total_start_time_string = datetime.datetime.fromtimestamp(total_start_time).strftime('%Y-%m-%d %H:%M:%S')
- utils.print_flush("total start time\t: %s" % (total_start_time_string), script_name = script_name)
- ## Input Data Loading
- saving_start_time = time.time()
- utils.exit_if_not_exist_file(file_path = input_path, script_name = script_name)
- utils.print_flush("Loading input variable X from ( %s )." % (input_path), script_name = script_name)
- X = utils.load_hdf5storage_data(file_path = input_path, variable_name = "X")
- m = np.shape(X)[0]
- n = np.shape(X)[1]
- if np.shape(np.shape(X))[0] != 2:
- 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)
- sys.exit(1)
- utils.print_verbose("Loaded input variable X of shape [%i, %i] from ( %s )." % (m, n, input_path), script_name = script_name)
- #calculate l
- l = np.amin([n, np.amax([target_rank + oversampling, target_rank + 10])])
- utils.print_flush(string_variable = "l: ( %i )" % (l), script_name = script_name)
- outdir = restart.get_output_directory(
- output_parent_dir = output_parent_dir,
- nmf_method = nmf_method,
- target_rank = target_rank,
- tol = tol,
- max_iter = max_epoch,
- update_method = multiplicative_update_method,
- ORTHONORMAL_FLAG = ORTHONORMAL_FLAG,
- init_method = init_method,
- MEM_FLAG = MEM_FLAG,
- sampling_method = sampling_method,
- 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
- eta = eta,
- batch_size = batch_size,
- X_compression_size =Q_X_compression_size,
- l = l,
- ZERO_REMOVAL_FLAG = ZERO_REMOVAL_FLAG,
- SMALL_VALUE_RESET_INIT_FLAG = SMALL_VALUE_RESET_INIT_FLAG,
- SMALL_VALUE_RESET_FLAG = SMALL_VALUE_RESET_FLAG,
- min_reset_value = min_reset_value,
- EPSILON = EPSILON,
- iteration_per_batch = iteration_per_batch,
- power_iteration_compression = power_iteration,
- COMPRESSION_PER_BATCH_FLAG = COMPRESSION_PER_BATCH_FLAG,
- script_name = script_name
- )
- output_basename_prefix = "%s_%s" % (nmf_method, update_meth)
- #define output directory
- #To-Do: replace with restart.py output directory function
- # 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))
- # if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize" or multiplicative_update_method == "quadratic" or multiplicative_update_method == "quadraticOrthonormal":
- # outdir = os.path.join(outdir, "rhoInit%0.2E" % (rho0))
- # if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
- # outdir = os.path.join(outdir, "eta%0.2E" % (eta))
- # if sampling_method == "dppggaussian":
- # outdir = os.path.join(outdir, "sigma%0.2E" % (sigma))
- # output_path = os.path.join(outdir, "%s.mat" % (output_basename_prefix))
- output_path = restart.get_output_path(
- output_parent_dir = output_parent_dir,
- nmf_method = nmf_method,
- target_rank = target_rank,
- tol = tol,
- max_iter = max_epoch,
- update_method = multiplicative_update_method,
- ORTHONORMAL_FLAG = ORTHONORMAL_FLAG,
- init_method = init_method,
- MEM_FLAG = MEM_FLAG,
- basename_prefix = output_basename_prefix,
- sampling_method = sampling_method,
- 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
- eta = eta,
- batch_size = batch_size,
- X_compression_size =Q_X_compression_size,
- l = l,
- ZERO_REMOVAL_FLAG = ZERO_REMOVAL_FLAG,
- SMALL_VALUE_RESET_INIT_FLAG = SMALL_VALUE_RESET_INIT_FLAG,
- SMALL_VALUE_RESET_FLAG = SMALL_VALUE_RESET_FLAG,
- min_reset_value = min_reset_value,
- EPSILON = EPSILON,
- iteration_per_batch = iteration_per_batch,
- power_iteration_compression = power_iteration,
- COMPRESSION_PER_BATCH_FLAG = COMPRESSION_PER_BATCH_FLAG,
- script_name = script_name
- )
- utils.print_flush(string_variable = "outdir: ( %s )" % (outdir), script_name = script_name)
- utils.print_flush(string_variable = "output_path: ( %s )" % (output_path), script_name = script_name)
- #output path
- restart_path = os.path.join(outdir, "%s_restart.mat" % (output_basename_prefix))
- restart_path_old = os.path.join(outdir, "%s_restart_old.mat" % (output_basename_prefix))
- initialization_path = os.path.join(outdir, "%s_initialization.mat" % (output_basename_prefix))
- sampling_intermediate_path = os.path.join(outdir, "%s_sampling.mat" % (output_basename_prefix))
- utils.print_flush(string_variable = "restart_path: ( %s )" % (restart_path), script_name = script_name)
- utils.print_flush(string_variable = "initialization_path: ( %s )" % (initialization_path), script_name = script_name)
- utils.print_flush(string_variable = "sampling_intermediate_path: ( %s )" % (sampling_intermediate_path), script_name = script_name)
- #sanity check: does parent output directory exist?
- utils.print_verbose("Checking if output directory ( %s ) exists." % (output_parent_dir), script_name = script_name, VERBOSE_FLAG = VERBOSE_FLAG)
- utils.exit_if_not_exist_dir(dir_path = output_parent_dir, script_name = script_name)
- # outdir = os.path.join(outdir, "Q_X_compression_size%i" % (Q_X_compression_size), "l%i" % (l), "powerIter%0.2E" % (power_iteration))
- # utils.print_flush(string_variable = "outdir: ( %s )" % (outdir), script_name = script_name)
- #sanity check: does output exist?
- if not os.path.isdir(outdir):
- os.makedirs(outdir)
- utils.exit_if_exist_file(file_path = output_path, script_name = script_name)
- #load demographics data if not uniform sampling
- if (sampling_method == "dpplinear") or (sampling_method == "dppgaussian"):
- utils.print_flush(string_variable = "demographic_data_path: ( %s )" %(demographic_data_path), script_name = script_name)
- utils.exit_if_not_exist_file(demographic_data_path)
- demographic_data = pd.read_csv(demographic_data_path)
- #sanity check: does k <= n and k <= m
- 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)
- if target_rank > m:
- 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)
- sys.exit(1)
- if target_rank > n:
- 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)
- sys.exit(1)
- saving_elapsed_time = saving_elapsed_time + (time.time() - saving_start_time)
- del saving_start_time
- ## Initializations
- #To-Do: load from initialization file if it exists instead of always initializing from scratch
- if os.path.isfile(initialization_path):
- saving_start_time = time.time()
- utils.print_verbose("Loading intialization data from ( %s )." % (initialization_path), script_name = script_name, VERBOSE_FLAG = VERBOSE_FLAG)
- w0 = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "w0", script_name = script_name)
- h0 = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "w0", script_name = script_name)
- saving_elapsed_time = saving_elapsed_time + ( time.time() - saving_start_time )
- del saving_start_time
- else:
- initialization_start_time = time.time()
- utils.print_verbose("Initializing W and H with %s intialization." % (init_method), script_name = script_name, VERBOSE_FLAG = VERBOSE_FLAG)
- if (init_method == "nndsvd") or (init_method == "nndsvda") or (init_method == "nndsvdar") or (init_method == "random"):
- w0, h0 = initialize_nmf._initialize_nmf(X, target_rank, init=init_method, eps=1e-6, random_state=None)
- else:
- utils.print_flush("ERROR - Unknown init_method (%s). Exiting." % (init_method), script_name = script_name)
- initialization_elapsed_time = initialization_elapsed_time + (time.time() - initialization_start_time)
- del initialization_start_time
- #initialize rest of variables (XX, W_old, diffW, etc.)
- initialization_start_time = time.time()
- #XX^T if not using MEM mode
- utils.print_verbose("Calculating XX^T to store in memory.", script_name = script_name)
- if (MEM_FLAG == False):
- XX = np.matmul(X, np.transpose(X))
- #initialize W to w0
- W = w0
- W_old = W
- #initialize diffW for print purposes and for stopping criterion
- diffW = 0.0
- initialization_elapsed_time = initialization_elapsed_time + (time.time() - initialization_start_time)
- del initialization_start_time
- ## Sampling
- #load idx_sample if inititialization path exists
- if os.path.isfile(initialization_path):
- saving_start_time = time.time()
- num_batch = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "num_batch", script_name = script_name)
- idx_sample = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "idx_sample", script_name = script_name)
- saving_elapsed_time = saving_elapsed_time + ( time.time() - saving_start_time )
- del saving_start_time
- #if initialization path does not exist, sample from scratch
- else:
- sampling_start_time = time.time()
- utils.print_verbose("Calculating num_batch based on n ( %i ) and batch_size ( %i )" % (n, batch_size), script_name = script_name)
- num_batch = int(np.ceil(np.divide(float(n), float(batch_size)))) #we need this value as int
- if (num_batch != np.divide(float(n), float(batch_size))):
- 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)
- #generate indices for all iterations/epochs
- utils.print_flush("Sampling indices of subjects using %s sampling method" % (sampling_method), script_name = script_name)
- if (batch_size == n):
- 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)
- idx_sample = -1 * np.ones(shape = (n, 1, max_epoch))
- for epoch in np.arange(max_epoch):
- idx_sample[:,:,epoch] = np.reshape(np.arange(n), newshape = (n, 1))
- print("%s: idx_sample = " % (script_name), flush=True)
- print(idx_sample, flush=True)
- elif (sampling_method == "uniform"):
- 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)
- elif (sampling_method == "dpplinear"):
- #save every 100 epochs
- 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 )
- elif (sampling_method == "dppgaussian"):
- 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 )
- else:
- utils.print_flush("Unknown sampling method (%s). Exiting.", script_name = script_name)
- sys.exit(1)
- utils.print_flush("Finished sampling indices of subjects using %s sampling method" % (sampling_method), script_name = script_name)
- sampling_elapsed_time = sampling_elapsed_time + (time.time() - sampling_start_time)
- del sampling_start_time
- #calculate Q for compression
- if not COMPRESSION_PER_BATCH_FLAG:
- if os.path.isfile(initialization_path):
- saving_start_time = time.time()
- Q = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "Q", script_name = script_name)
- saving_elapsed_time = saving_elapsed_time + ( time.time() - saving_start_time )
- del saving_start_time
- else:
- compression_start_time = time.time()
- 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)
- if VERBOSE_FLAG:
- print("Q.shape = ", flush=True)
- print(Q.shape, flush=True) #should be l by l?
- compression_elapsed_time = compression_elapsed_time + (time.time() - compression_start_time)
- del compression_start_time
- #sanity check: if Q_X_compression_size > batch_size, select first batch_size of Q_X_compression_size of Q to use for compression
- compression_start_time = time.time()
- if Q_X_compression_size < batch_size:
- 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)
- sys.exit(1)
- elif Q_X_compression_size == batch_size:
- Q = Q_full.copy()
- else:
- Q = Q_full[0:batch_size, :]
- if VERBOSE_FLAG:
- 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)
- del Q_full
- compression_elapsed_time = compression_elapsed_time + (time.time() - compression_start_time)
- del compression_start_time
- 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)
- utils.print_verbose("X.shape = [%i, %i]" % (m, n) , script_name = script_name)
- utils.print_verbose("X.min max = [%0.15f, %0.15f]" % (np.amin(X, axis = None), np.amax(X, axis = None) ) , script_name = script_name)
- utils.print_verbose("w0.shape = [%i, %i]" % (np.shape(w0)[0], np.shape(w0)[1]) , script_name = script_name)
- utils.print_verbose("w0.min max = [%0.15f, %0.15f]" % (np.amin(w0, axis = None), np.amax(w0, axis = None) ) , script_name = script_name)
- utils.print_verbose("h0.shape = [%i, %i]" % (np.shape(h0)[0], np.shape(h0)[1]) , script_name = script_name)
- utils.print_verbose("h0.min max = [%0.15f, %0.15f]" % (np.amin(h0, axis = None), np.amax(h0, axis = None) ) , script_name = script_name)
- #print information about batch sampling
- utils.print_flush("batch_size: ( %i )" % (batch_size), script_name = script_name)
- utils.print_verbose("batch_size (number of subjects per batch): %i" % (batch_size), script_name = script_name)
- utils.print_flush("num_batch: ( %i )" % (num_batch), script_name = script_name)
- utils.print_verbose("num_batch (number of batches in full X): %i" % (num_batch), script_name = script_name)
- utils.print_flush("max_epoch: ( %i )" % (max_epoch), script_name = script_name)
- utils.print_flush("iteration_per_batch: ( %i )" % (iteration_per_batch), script_name = script_name)
- #initalize arrays to store values during the for loop of updates
- epoch_start = 0
- statistics_start_time = time.time()
- epoch_array = np.arange(start = 0, stop = max_epoch, step = 1)
- batch_loss_sum_per_epoch_array = np.zeros(shape = epoch_array.shape)
- full_loss_per_epoch_array = np.zeros(shape = epoch_array.shape)
- batch_loss_per_iteration_array = np.zeros(shape = (max_epoch, num_batch))
- full_loss_per_iteration_array = np.zeros(shape = (max_epoch, num_batch))
- sparsity_per_epoch_array = np.zeros(shape = epoch_array.shape)
- sparsity_per_iteration_array = np.zeros(shape = (max_epoch, num_batch))
- if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
- trXtX = np.power(np.linalg.norm(X, 'fro'), 2)
- obj = utils.get_objective_function(X = X, W = W, trXtX = trXtX, OPNMF = True)
- obj_old = obj
- objective_function = 0.0 #update to value if VERBOSE_FLAG is on
- statistics_elapsed_time = statistics_elapsed_time + (time.time() - statistics_start_time)
- del statistics_start_time
- #to facilitate |X-WH|_F^2 objective calculation, we want to store tr(X'X) in the memory and reuse it
- initialization_start_time = time.time()
- if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
- trXtX = np.power(np.linalg.norm(X, 'fro'), 2)
- 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
- obj_old = obj
- initialization_elapsed_time = initialization_elapsed_time + (time.time() - initialization_start_time)
- del initialization_start_time
- ## load intermediate if it exists
- if os.path.isfile(initialization_path):
- saving_start_time = time.time()
- utils.print_flush(string_variable = "Loading initialization from file ( %s )." % (initialization_path), script_name = script_name)
- total_elapsed_time = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "total_elapsed_time", script_name = script_name)
- saving_elapsed_time = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "saving_elapsed_time", script_name = script_name)
- initialization_elapsed_time = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "initialization_elapsed_time", script_name = script_name)
- statistics_elapsed_time = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "statistics_elapsed_time", script_name = script_name)
- compression_elapsed_time = utils.load_hdf5storage_data(file_path = initialization_path, variable_name = "compression_elapsed_time", script_name = script_name)
- saving_elapsed_time = saving_elapsed_time + (time.time() - saving_start_time)
- del saving_start_time
- else:
- saving_start_time = time.time()
- mdict = {
- "w0": w0,
- "h0": h0,
- "num_batch": num_batch,
- "idx_sample": idx_sample,
- "total_elapsed_time": total_elapsed_time,
- "saving_elapsed_time": saving_elapsed_time,
- "initialization_elapsed_time": initialization_elapsed_time,
- "statistics_elapsed_time": statistics_elapsed_time,
- "compression_elapsed_time": compression_elapsed_time
- }
- if not COMPRESSION_PER_BATCH_FLAG:
- mdict_temp = {"Q": Q}
- mdict.update(mdict_temp)
- del mdict_temp
- utils.save_intermediate_output(output_path = initialization_path, mdict = mdict, OVERWRITE = False)
- saving_elapsed_time = saving_elapsed_time + (time.time() - saving_start_time)
- del saving_start_time
- ## load from restart point
- if os.path.isfile(restart_path):
- utils.print_flush(string_variable = "Loading restart file from file ( %s )." % (restart_path), script_name = script_name)
- epoch_start = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "epoch")
- W = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "W")
- W_old = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "W_old")
- diffW = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "diffW")
- total_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "total_elapsed_time")
- initialization_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "initialization_elapsed_time")
- saving_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "saving_elapsed_time")
- statistics_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "statistics_elapsed_time")
- update_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "update_elapsed_time")
- sampling_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "sampling_elapsed_time")
- compression_elapsed_time = utils.load_hdf5storage_data(file_path = restart_path, variable_name = "compression_elapsed_time")
- utils.print_flush(string_variable = "Loading from restart: epoch ( %i ) and elapsed time ( %s seconds )." % (epoch_start, total_elapsed_time), script_name = script_name)
- #start for loop
- for epoch in np.arange(start = epoch_start, stop = epoch_array[-1], step = 1):
- #restart happens at epoch level
- if np.mod(epoch, restart_step) == 0:
- saving_start_time = time.time()
- # if epoch != 0:
- # #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)
- # time.sleep(0.5) #sleep for 0.5 seconds
- # shutil.move(src = restart_path, dst = restart_path_old)
- # time.sleep(0.5) #sleep for 0.5 seconds
- utils.print_verbose(string_variable = "epoch %i / %i: saving intermediate save point for restarting." % (epoch, max_epoch), script_name = script_name)
- mdict = {
- "W": W,
- "W_old": W_old,
- "diffW": diffW,
- "epoch": epoch,
- "initialization_elapsed_time": initialization_elapsed_time,
- "saving_elapsed_time": saving_elapsed_time,
- "statistics_elapsed_time": statistics_elapsed_time,
- "update_elapsed_time": update_elapsed_time,
- "sampling_elapsed_time": sampling_elapsed_time,
- "compression_elapsed_time": compression_elapsed_time,
- "total_elapsed_time": total_elapsed_time + (time.time() - total_start_time)
- }
- utils.save_intermediate_output(output_path = restart_path, mdict = mdict, OVERWRITE = True)
- # if not os.path.isfile(restart_path_old): #must be missing restart_path_old because epoch == 0 or it got deleted
- # utils.save_intermediate_output(output_path = restart_path_old, mdict = mdict, OVERWRITE = False)
- restart_path_epoch = os.path.join(outdir, "%s_restart_epoch%05d.mat" % (output_basename_prefix, epoch))
- utils.print_verbose("Copying ( %s ) to ( %s ) as additional backup." % (restart_path, restart_path_epoch), script_name = script_name)
- shutil.copy(src = restart_path, dst = restart_path_epoch)
- restart_path_epoch_to_remove = os.path.join(outdir, "%s_restart_epoch%05d.mat" % (output_basename_prefix, epoch - 2 * restart_step))
- if os.path.isfile(restart_path_epoch_to_remove):
- utils.print_verbose("Removing very old restart file ( %s )." % (restart_path_epoch_to_remove), script_name = script_name)
- os.unlink(restart_path_epoch_to_remove)
- del restart_path_epoch_to_remove
- del restart_path_epoch
- del mdict
- saving_elapsed_time = saving_elapsed_time + (time.time() - saving_start_time)
- del saving_start_time
- 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
- ## Batch Sampling: select the batch of subjects from full X for this iteration
- sampling_start_time = time.time()
- 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)
- idx_sample_selected = (idx_sample[:, int(batch_index), int(epoch)]).astype(int)
- X_sampled = X[:, idx_sample_selected]
- del idx_sample_selected
- sampling_elapsed_time = sampling_elapsed_time + (time.time() - sampling_start_time)
- del sampling_start_time
- #compress X_sampled to X_sampled
- if VERBOSE_FLAG:
- print("X_sampled.shape = ", flush=True)
- print(X_sampled.shape, flush=True)
- #if COMPRESSION_PER_BATCH_FLAG is on, generate Q FROM batch X_sampled here to compress X_sampled
- compression_start_time = time.time()
- if COMPRESSION_PER_BATCH_FLAG:
- if Q_X_compression_size < batch_size:
- 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)
- Q_X_compression_size = batch_size
- 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)
- if VERBOSE_FLAG:
- print("Q_full.shape = ", flush=True)
- print(Q_full.shape, flush=True) #should be l by l?
- #if Q_X_compression_size > batch_size, select first batch_size of Q_X_compression_size of Q to use for compression
- if Q_X_compression_size < batch_size:
- 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)
- sys.exit(1)
- elif Q_X_compression_size == batch_size:
- Q = Q_full.copy()
- else:
- Q = Q_full[0:batch_size, :]
- if VERBOSE_FLAG:
- 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)
- #Q_full is no longer needed now that we have final Q to use
- del Q_full
- if epoch == 0 and batch_index == 0:
- 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)
- # X_sampled_compressed = np.matmul(Q.T, X_sampled) #if compressing m -> l
- X_sampled_compressed = np.matmul(X_sampled, Q) #if compressing n -> l
- if VERBOSE_FLAG:
- print("X_sampled_compressed.shape = ",flush=True)
- print(X_sampled_compressed.shape, flush=True)
- compression_elapsed_time = compression_elapsed_time + ( time.time() - compression_start_time )
- del compression_start_time
- for iteration in np.arange(iteration_per_batch): #one iteration per batch
- saving_start_time = time.time()
- #save small file for plotting progress after the run finishes
- if np.mod(epoch, save_step) == 0 and ( (batch_index == 0) and (iteration == 0) ):
- current_elapsed_time = (time.time() - total_start_time) + (total_elapsed_time)
- intermediate_debug_path = os.path.join(outdir, "%s_intermediate_epoch%05d.mat" % (output_basename_prefix, epoch))
- 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)
- mdict = {
- "W": W,
- "elapsed_time": current_elapsed_time,
- "epoch": epoch,
- "iteration": iteration,
- "batch_index": batch_index
- }
- hdf5storage.savemat(intermediate_debug_path, mdict = mdict)
- del mdict, intermediate_debug_path
- #print progress to console
- if np.mod(epoch, print_step) == 0 and batch_index == 0 and iteration == 0:
- current_elapsed_time = (time.time() - total_start_time) + (total_elapsed_time)
- remaining_time_string = utils.get_remaining_time_string(iteration = epoch, max_iter = max_epoch, elapsed_time = (time.time() - total_start_time) + total_elapsed_time )
- 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)
- del current_elapsed_time, remaining_time_string
- saving_elapsed_time = saving_elapsed_time + (time.time() - saving_start_time)
- del saving_start_time
- utils.print_debug(string_variable = "Storing old W", script_name = script_name)
- W_old = W
- if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
- utils.print_debug(string_variable = "Storing old objective function = %0.5E" % (obj), script_name = script_name)
- obj_old = obj
- #multiplicative update for W
- update_start_time = time.time()
- utils.print_debug(string_variable = "Update with %s multiplicative update rule on W" % (multiplicative_update_method), script_name = script_name)
- if (multiplicative_update_method == "original") or (multiplicative_update_method == "normalize"):
- #multiplicative update rule - modified version that does not require XX = X*X' ready and available
- W = opnmf_update_rule.multiplicative_update(X = X_sampled_compressed, W = W, OPNMF = True, MEM = True, rho = 1, DEBUG = False, script_name = script_name)
- elif (multiplicative_update_method == "decoupled"):
- W = opnmf_update_rule.multiplicative_update_decoupled(W = W, X = X_sampled_compressed, script_name = script_name, SANITY_CHECK_FLAG = False)
- elif (multiplicative_update_method == "adaptive") or (multiplicative_update_method == "adaptiveNormalize"):
- 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)
- elif (update_meth == "quadratic"):
- W = opnmf_update_rule.multiplicative_update_quadratic(W = W, X = X_sampled_compressed, OPNMF = False, MEM = True, rho = rho0, EPSILON = EPSILON)
- elif (update_meth == "quadraticOrthonormal"):
- W = opnmf_update_rule.multiplicative_update_quadratic(W = W, X = X_sampled_compressed, OPNMF = True, MEM = True, rho = rho0, EPSILON = EPSILON)
- else:
- print("%s: multiplicative_update_method (%s) must be either original, normalize, decoupled, adaptive, adaptiveNormalize, quadratic, quadraticOrthonormal. Exiting." % (script_name, multiplicative_update_method), flush=True)
- sys.exit(1)
- #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)
- utils.print_debug(string_variable = "Resetting small values of W", script_name = script_name)
- W = opnmf_update_rule.reset_small_value(W, min_reset_value = reset_small_value_threshold)
- #normalize W following mulitplicative update to stabilize W; only if not using normalization mulitplicative update version, not for convergent (constant root) or adaptive update
- if multiplicative_update_method == "normalize" or multiplicative_update_method == "adaptiveNormalize":
- utils.print_debug(string_variable = "Normalizing W", script_name = script_name)
- W = opnmf_update_rule.normalize_W(W = W)
- if multiplicative_update_method == "adaptive" or multiplicative_update_method == "adaptiveNormalize":
- utils.print_debug(string_variable = "updating rho ( %0.5E ) for %s multiplicative update method." % (rho, multiplicative_update_method), flush=True)
- bOK = False
- while not bOK:
- 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
- if rho != rho0 and obj > obj_old:
- bOK = False
- rho = rho0
- W = W_old
- obj = obj_old
- else:
- bOK = True
- rho = rho + eta
- update_elapsed_time = update_elapsed_time + (time.time() - update_start_time)
- del update_start_time
- ## stopping criterion: diffW for normalization or convergent update using diffW
- update_start_time = time.time()
- diffW = np.linalg.norm(W_old - W, ord = 'fro') / np.linalg.norm(W_old, ord = 'fro')
- utils.print_debug(string_variable = "diffW = %0.5E" % (diffW), script_name = script_name)
- if multiplicative_update_method == "constant" or multiplicative_update_method == "normalize":
- if diffW < tol:
- 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)
- update_elapsed_time = update_elapsed_time + (time.time() - update_start_time)
- del update_start_time
- ## Calculate H
- update_start_time = time.time()
- utils.print_flush(string_variable = "calculating final H", script_name = script_name)
- H = np.matmul(np.transpose(W), X)
- update_elapsed_time = update_elapsed_time + (time.time() - update_start_time)
- del update_start_time
- total_end_time = time.time()
- total_end_time_string = datetime.datetime.fromtimestamp(total_end_time).strftime('%Y-%m-%d %H:%M:%S')
- utils.print_flush(string_variable = "total start time\t: %s" % (total_start_time_string), script_name = script_name)
- utils.print_flush(string_variable = "total end time\t: %s" % (total_end_time_string), script_name = script_name)
- del total_start_time_string, total_end_time_string
- total_elapsed_time = total_elapsed_time + (total_end_time - total_start_time)
- #save final output
- saving_start_time = time.time()
- utils.print_flush(string_variable = "saving final output (%s)" % (output_path), script_name = script_name)
- mdict = {
- "X": X,
- "W": W,
- "H": H,
- "Q": Q,
- "batch_loss_per_iteration_array": batch_loss_per_iteration_array,
- "full_loss_per_iteration_array": full_loss_per_iteration_array,
- "sparsity_per_iteration_array": sparsity_per_iteration_array,
- "batch_loss_sum_per_epoch_array": batch_loss_sum_per_epoch_array,
- "full_loss_per_epoch_array": full_loss_per_epoch_array,
- "sparsity_per_epoch_array": sparsity_per_epoch_array,
- "total_elapsed_time": total_elapsed_time,
- "initialization_elapsed_time": initialization_elapsed_time,
- "saving_elapsed_time": saving_elapsed_time,
- "statistics_elapsed_time": statistics_elapsed_time,
- "update_elapsed_time": update_elapsed_time,
- "sampling_elapsed_time": sampling_elapsed_time,
- "compression_elapsed_time": compression_elapsed_time
- }
- utils.save_intermediate_output(output_path = output_path, mdict = mdict)
- del mdict
- saving_elapsed_time = saving_elapsed_time + (time.time() - saving_start_time)
- del saving_start_time
- utils.print_flush(string_variable = "Total optimization elapsed time: %f seconds" % (total_elapsed_time))
- utils.print_flush(string_variable = "Saving/Restarting/IO: %f seconds" % (saving_elapsed_time))
- utils.print_flush(string_variable = "Multiplicative Updates: %f seconds" % (update_elapsed_time))
- utils.print_flush(string_variable = "Initialization: %f seconds" % (initialization_elapsed_time))
- utils.print_flush(string_variable = "Statistics Calculations: %f seconds" % (statistics_elapsed_time))
- utils.print_flush(string_variable = "Sampling: %f seconds" % (sampling_elapsed_time))
- utils.print_flush(string_variable = "Compression: %f seconds" % (compression_elapsed_time))
- sys.exit(0)
csopnmf.py at commit 7b40455, under Apache-2.0 · at the source
Overview
- Mallinckrodt Institute of Radiology, Washington University in St. Louis, St. Louis, MO, United States
- Institute for Informatics, Data Science & Biostatistics (I, 2, DB), Washington University in St. Louis, St. Louis, MO, United States
- Department of Radiology, Perelman School of Medicine, University of Pennsylvania, Philadelphia, PA, United States
- Center for AI and Data Science for Integrated Diagnostics (AI2D), Perelman School of Medicine, University of Pennsylvania, Philadelphia, PA, United States
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
7b404557b3c78761e9c67b96c2ed8da5e36ce44d, 14 September 2026Availability: 1 check, the latest on 26 September 2026: the link answers
- 26 September 2026: the link answers
9 files
- Python/
csopnmf.py , Python, 1,469 lines, 1 match - Python/
initialize_nmf.py , Python, 184 lines, 1 match - Python/
opnmf.py , Python, 526 lines - Python/
opnmf_update_rule.py , Python, 381 lines - Python/
restart.py , Python, 191 lines - Python/
sopnmf.py , Python, 905 lines - Python/
utils.py , Python, 670 lines - LICENSE, License, 201 lines
- README.md, Text, 142 lines
The paper's code and data availability statement is in the Data section.
Tracing map
Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.
What the map holds:
- 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
- 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://
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://
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/
url = {https://
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/
VL - 4
SP - IMAG.a.1355
SN - 2837-6056
PB - MIT Press
DO - 10.1162/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1162/
"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":
"volume": "4",
"page": "IMAG.a.1355",
"DOI": "10.1162/
"PMID": "42719766",
"PMCID": "PMC13556794",
"ISSN": "2837-6056",
"publisher": "MIT Press",
"URL": "https://
"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 communicationsIn 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 communicationsIn 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 brainJournal: 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 brainJournal: 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 biologyIn 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 communicationsIn 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 AmericaIn 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 communicationsIn 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 dataIn 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: eLifeIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 7 scripts, and 2 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:862f204fadf210d1…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
