From bede1f60460f413abb060bbe02cf4007f4a040c3 Mon Sep 17 00:00:00 2001 From: tqlong Date: Fri, 28 May 2010 14:12:12 +0000 Subject: [PATCH] anmf --- .../contrib/tqlong/affineNMF/ANMF_test.cc | 55 +++++++++++++++++-- fastlib/trunk/contrib/tqlong/affineNMF/nmf.cc | 46 +++++++++++++++- fastlib/trunk/contrib/tqlong/affineNMF/nmf.h | 4 +- 3 files changed, 94 insertions(+), 11 deletions(-) diff --git a/fastlib/trunk/contrib/tqlong/affineNMF/ANMF_test.cc b/fastlib/trunk/contrib/tqlong/affineNMF/ANMF_test.cc index 08d90f8479..12b532cb81 100644 --- a/fastlib/trunk/contrib/tqlong/affineNMF/ANMF_test.cc +++ b/fastlib/trunk/contrib/tqlong/affineNMF/ANMF_test.cc @@ -4,9 +4,9 @@ const fx_entry_doc anmf_entries[] = { {"i1", FX_PARAM, FX_STR, NULL, - " HMM type : discrete | gaussian | mixture.\n"}, + " input file 1.\n"}, {"i2", FX_PARAM, FX_STR, NULL, - " A file containing HMM transition.\n"}, + " input file 2.\n"}, /* {"fileE", FX_REQUIRED, FX_STR, NULL, " A file containing HMM emission.\n"}, @@ -33,17 +33,60 @@ const fx_module_doc anmf_doc = { "This is a program generating sequences from HMM models.\n" }; +void InitRandom01(index_t n_rows, index_t n_cols, Matrix* A_) { + Matrix& A = *A_; + A.Init(n_rows, n_cols); + for (index_t i = 0; i < n_rows; i++) + for (index_t j = 0; j < n_cols; j++) + A.ref(i, j) = math::Random(0.1,1.0000); +} + +void nmf_run(const Matrix& V, index_t rank, + Matrix* W_, Matrix* H_) { + Matrix Winit, Hinit; + InitRandom01(V.n_rows(), rank, &Winit); + InitRandom01(rank, V.n_cols(), &Hinit); + + nmf(V, Winit, Hinit, 10, W_, H_); +} + int main(int argc, char* argv[]) { fx_module* root = fx_init(argc, argv, &anmf_doc); - Matrix I1, I2; - data::Load(fx_param_str(root, "i1", "i1"), &I1); - data::Load(fx_param_str(root, "i2", "i2"), &I2); + const char* f1 = fx_param_str(root, "i1", "i1"); + const char* f2 = fx_param_str(root, "i2", "i2"); + // Test registration + Matrix I1, I2; + data::Load(f1, &I1); + data::Load(f2, &I2); Vector m; projective_register(I1, I2, &m); - ot::Print(m); + + /* Test nmf */ + Matrix V; + + V.Init(400,5); + for (int i = 0; i < 5; i++) { + Matrix X; + char fn[100]; + sprintf(fn, "im%d", i+1); + data::Load(fn, &X); + for (int j = 0; j < 400; j++) { + V.ref(j,i) = X.get(j%20, j/20)/255; + } + } + + //data::Load("V", &V); + prepare_for_nmf(V); + printf("size(V) = %d x %d", V.n_rows(), V.n_cols()); + Matrix W, H; + nmf_run(V, 2, &W, &H); + data::Save("basis", W); + printf("size(W) = %d x %d", W.n_rows(), W.n_cols()); + data::Save("weight", H); + printf("size(H) = %d x %d", H.n_rows(), H.n_cols()); fx_done(root); return 0; } diff --git a/fastlib/trunk/contrib/tqlong/affineNMF/nmf.cc b/fastlib/trunk/contrib/tqlong/affineNMF/nmf.cc index 3c4d6314ca..59eed4b293 100644 --- a/fastlib/trunk/contrib/tqlong/affineNMF/nmf.cc +++ b/fastlib/trunk/contrib/tqlong/affineNMF/nmf.cc @@ -2,7 +2,47 @@ #include #include "nmf.h" -void nmf(const Matrix& V, const Matrix& Winit, const Matrix& Hinit, - double tol, double timelimit, index_t maxiter, - Matrix* W_, Matrix* H_) { +void nmf(const Matrix& V, const Matrix& Winit, const Matrix& Hinit, + index_t maxiter, Matrix* W_, Matrix* H_) { + Matrix& W = *W_; + Matrix& H = *H_; + + W.Copy(Winit); + H.Copy(Hinit); + + for (index_t iter = 0; iter < maxiter; iter++) { + //ot::Print(H); + // update H + Matrix WtV, WtW, WtWH; + la::MulTransAInit(W, V, &WtV); + la::MulTransAInit(W, W, &WtW); + la::MulInit(WtW, H, &WtWH); + //ot::Print(WtWH, "WtWH"); + //ot::Print(WtV, "WtV"); + for (index_t i = 0; i < H.n_rows(); i++) + for (index_t j = 0; j < H.n_cols(); j++) { + if (WtWH.get(i,j) < 1e-10) + printf("WtWh bad"); + H.ref(i, j) *= WtV.get(i, j) / WtWH.get(i, j); + } + + // update W + Matrix VHt, HHt, WHHt; + la::MulTransBInit(V, H, &VHt); + la::MulTransBInit(H, H, &HHt); + la::MulInit(W, HHt, &WHHt); + for (index_t i = 0; i < W.n_rows(); i++) + for (index_t j = 0; j < W.n_cols(); j++) { + if (WHHt.get(i, j) < 1e-10) + printf("WHHt bad."); + W.ref(i, j) *= VHt.get(i, j) / WHHt.get(i, j); + } + } +} + + +void prepare_for_nmf(Matrix& V) { + for (index_t i = 0; i < V.n_rows(); i++) + for (index_t j = 0; j < V.n_cols(); j++) + if (V.get(i, j) < 1e-8) V.ref(i, j) = 1e-4; } diff --git a/fastlib/trunk/contrib/tqlong/affineNMF/nmf.h b/fastlib/trunk/contrib/tqlong/affineNMF/nmf.h index f55f8dd9a7..f7572386a4 100644 --- a/fastlib/trunk/contrib/tqlong/affineNMF/nmf.h +++ b/fastlib/trunk/contrib/tqlong/affineNMF/nmf.h @@ -4,5 +4,5 @@ **/ void nmf(const Matrix& V, const Matrix& Winit, const Matrix& Hinit, - double tol, double timelimit, index_t maxiter, - Matrix* W_, Matrix* H_); + index_t maxiter, Matrix* W_, Matrix* H_); +void prepare_for_nmf(Matrix& V);