From cc8a692a9a377be45c59e02e66294ddaee10c5fc Mon Sep 17 00:00:00 2001 From: Dongryeol Lee Date: Mon, 19 May 2008 17:29:29 +0000 Subject: [PATCH] Added the basis training for the local expansion of the leaf nodes. --- ...atrix_factorized_farfield_expansion_impl.h | 2 +- .../matrix_factorized_local_expansion.h | 8 +- .../matrix_factorized_local_expansion_impl.h | 76 +++++++++++++++++++ 3 files changed, 81 insertions(+), 5 deletions(-) diff --git a/fastlib2/mlpack/series_expansion/matrix_factorized_farfield_expansion_impl.h b/fastlib2/mlpack/series_expansion/matrix_factorized_farfield_expansion_impl.h index eea7ecf58c..09269030b4 100644 --- a/fastlib2/mlpack/series_expansion/matrix_factorized_farfield_expansion_impl.h +++ b/fastlib2/mlpack/series_expansion/matrix_factorized_farfield_expansion_impl.h @@ -28,7 +28,7 @@ void MatrixFactorizedFarFieldExpansion::AccumulateCoeffs // query samples taken from the stratification. Matrix sample_kernel_matrix; int num_reference_samples = (int) sqrt(end - begin); - int num_query_samples = (int) query_leaf_nodes->size(); + int num_query_samples = query_leaf_nodes->size(); // Allocate a temporary space for holding the indices of the // reference points, from which the outgoing skeleton will be diff --git a/fastlib2/mlpack/series_expansion/matrix_factorized_local_expansion.h b/fastlib2/mlpack/series_expansion/matrix_factorized_local_expansion.h index 0b68332296..24305f7684 100644 --- a/fastlib2/mlpack/series_expansion/matrix_factorized_local_expansion.h +++ b/fastlib2/mlpack/series_expansion/matrix_factorized_local_expansion.h @@ -27,10 +27,10 @@ class MatrixFactorizedLocalExpansion { */ Vector coeffs_; - /** @brief The incoming representation: the pseudo-distribution - * which is defined only for leaf nodes. + /** @brief The evaluation operator: the pseudo-distribution which is + * defined only for leaf nodes. */ - Vector *incoming_representation_; + Matrix *evaluation_operator_; /** @brief The query point indices that form the incoming skeleton, * the pseudo-points that represent the query point @@ -67,7 +67,7 @@ class MatrixFactorizedLocalExpansion { OT_DEF(MatrixFactorizedLocalExpansion) { OT_MY_OBJECT(coeffs_); - OT_PTR_NULLABLE(incoming_representation_); + OT_PTR_NULLABLE(evaluation_operator_); OT_MY_OBJECT(incoming_skeleton_); OT_MY_OBJECT(local_to_local_translation_begin_); OT_MY_OBJECT(local_to_local_translation_count_); diff --git a/fastlib2/mlpack/series_expansion/matrix_factorized_local_expansion_impl.h b/fastlib2/mlpack/series_expansion/matrix_factorized_local_expansion_impl.h index 9bb7e885a9..f1af6f1ff0 100644 --- a/fastlib2/mlpack/series_expansion/matrix_factorized_local_expansion_impl.h +++ b/fastlib2/mlpack/series_expansion/matrix_factorized_local_expansion_impl.h @@ -80,7 +80,83 @@ void MatrixFactorizedLocalExpansion::TrainBasisFunctions (const Matrix &query_set, int begin, int end, const Matrix *reference_set, const ArrayList *reference_leaf_nodes) { + // The sample kernel matrix is |Q| by S where |Q| is the number of + // query points in the query node and S is the number of reference + // samples taken from the stratification. + Matrix sample_kernel_matrix; + int num_reference_samples = reference_leaf_nodes->size(); + int num_query_samples = (int) sqrt(end - begin); + + // Allocate a temporary space for holding the indices of the query + // points, from which the incoming skeleton will be chosen. + ArrayList tmp_incoming_skeleton; + tmp_incoming_skeleton.Init(num_query_samples); + for(index_t q = 0; q < num_query_samples; q++) { + + // Choose a random query point and record its index. + index_t random_query_point_index = math::RandInt(begin, end); + tmp_incoming_skeleton[r] = random_query_point_index; + } + // Sort the chosen query indices and eliminate duplicates... + qsort(tmp_incoming_skeleton.begin(), tmp_incoming_skeleton.size(), + sizeof(index_t), &qsort_compar_); + remove_duplicates_in_sorted_array_(tmp_incoming_skeleton); + num_query_samples = tmp_incoming_skeleton.size(); + + // After determining the number of query samples to take, + // allocate the space for the sample kernel matrix to be computed. + sample_kernel_matrix.Init(num_query_samples, num_reference_samples); + for(index_t r = 0; r < num_reference_samples; r++) { + + // Choose a random reference point from the current reference strata... + index_t random_reference_point_index = + math::RandInt(((*reference_leaf_nodes)[r])->begin(), + ((*reference_leaf_nodes)[r])->end()); + const double *reference_point = + reference_set.GetColumnPtr(random_reference_point_index); + + for(index_t c = 0; c < num_query_samples; c++) { + + // The current query point + const double *query_point = + query_set->GetColumnPtr(tmp_incoming_skeleton[c]); + + // Compute the pairwise distance and the kernel value. + double squared_distance = + la::DistanceSqEuclidean(reference_set.n_rows(), reference_point, + query_point); + sample_kernel_matrix.set + (c, r, (ka_->kernel_).EvalUnnormOnSq(squared_distance)); + + } // end of iterating over each sample query strata... + } // end of iterating over each reference point... + + // CUR-decompose the sample kernel matrix. + Matrix c_mat, u_mat, r_mat; + ArrayList column_indices, row_indices; + CURDecomposition::Compute(sample_kernel_matrix, &c_mat, &u_mat, &r_mat, + &column_indices, &row_indices); + + // The incoming skeleton is constructed from the sampled rows in the + // matrix factorization. + incoming_skeleton_ = new ArrayList(); + incoming_skeleton_->Init(row_indices.size()); + for(index_t s = 0; s < row_indices.size(); s++) { + (*incoming_skeleton_)[s] = tmp_incoming_skeleton[row_incides[s]]; + } + + // Compute the evaluation operator, which is the product of the C + // and the U factor appropriately scaled by the row scaled R factor. + la::MulInit(c_mat, u_mat, evaluation_operator_); + for(index_t i = 0; i < r_mat.n_rows(); i++) { + double scaling_factor = + (sample_kernel_matrix.get(row_indices[i], 0) < DBL_EPSILON) ? + 0:r_mat.get(i, 0) / sample_kernel_matrix.get(row_indices[i], 0); + + la::Scale(evaluation_operator_->n_rows(), scaling_factor, + evaluation_operator_->GetColumnPtr(i)); + } } template