From 3d9546963d7c8c5f5dfb12a2df745f4996fd2ec5 Mon Sep 17 00:00:00 2001 From: Sameer Agarwal Date: Thu, 18 Apr 2013 14:54:55 -0700 Subject: [PATCH] Add the ability to query the Problem about parameter blocks. Change-Id: Ieda1aefa28e7a1d18fe6c8d1665882e4d9c274f2 --- include/ceres/problem.h | 13 +++++++++++++ internal/ceres/problem.cc | 12 ++++++++++++ internal/ceres/problem_impl.cc | 20 ++++++++++++++++++++ internal/ceres/problem_impl.h | 4 ++++ internal/ceres/problem_test.cc | 29 +++++++++++++++++++++++++++++ 5 files changed, 78 insertions(+) diff --git a/include/ceres/problem.h b/include/ceres/problem.h index 0a449cb99..707a8eb93 100644 --- a/include/ceres/problem.h +++ b/include/ceres/problem.h @@ -328,6 +328,19 @@ class Problem { // sizes of all of the residual blocks. int NumResiduals() const; + // The size of the parameter block. + int ParameterBlockSize(double* values) const; + + // The size of local parameterization for the parameter block. If + // there is no local parameterization associated with this parameter + // block, then ParmeterBlockLocalSize = ParameterBlockSize. + int ParameterBlockLocalSize(double* values) const; + + // Fills the passed parameter_blocks vector with pointers to the + // parameter blocks currently in the problem. After this call, + // parameter_block.size() == NumParameterBlocks. + void GetParameterBlocks(vector* parameter_blocks) const; + // Options struct to control Problem::Evaluate. struct EvaluateOptions { EvaluateOptions() diff --git a/internal/ceres/problem.cc b/internal/ceres/problem.cc index 43e78835b..b483932b2 100644 --- a/internal/ceres/problem.cc +++ b/internal/ceres/problem.cc @@ -206,4 +206,16 @@ int Problem::NumResiduals() const { return problem_impl_->NumResiduals(); } +int Problem::ParameterBlockSize(double* parameter_block) const { + return problem_impl_->ParameterBlockSize(parameter_block); +}; + +int Problem::ParameterBlockLocalSize(double* parameter_block) const { + return problem_impl_->ParameterBlockLocalSize(parameter_block); +}; + +void Problem::GetParameterBlocks(vector* parameter_blocks) const { + problem_impl_->GetParameterBlocks(parameter_blocks); +} + } // namespace ceres diff --git a/internal/ceres/problem_impl.cc b/internal/ceres/problem_impl.cc index f4615f947..34c378575 100644 --- a/internal/ceres/problem_impl.cc +++ b/internal/ceres/problem_impl.cc @@ -711,5 +711,25 @@ int ProblemImpl::NumResiduals() const { return program_->NumResiduals(); } +int ProblemImpl::ParameterBlockSize(double* parameter_block) const { + return FindParameterBlockOrDie(parameter_block_map_, parameter_block)->Size(); +}; + +int ProblemImpl::ParameterBlockLocalSize(double* parameter_block) const { + return FindParameterBlockOrDie(parameter_block_map_, + parameter_block)->LocalSize(); +}; + +void ProblemImpl::GetParameterBlocks(vector* parameter_blocks) const { + CHECK_NOTNULL(parameter_blocks); + parameter_blocks->resize(0); + for (ParameterMap::const_iterator it = parameter_block_map_.begin(); + it != parameter_block_map_.end(); + ++it) { + parameter_blocks->push_back(it->first); + } +} + + } // namespace internal } // namespace ceres diff --git a/internal/ceres/problem_impl.h b/internal/ceres/problem_impl.h index ccc315de6..260938964 100644 --- a/internal/ceres/problem_impl.h +++ b/internal/ceres/problem_impl.h @@ -139,6 +139,10 @@ class ProblemImpl { int NumResidualBlocks() const; int NumResiduals() const; + int ParameterBlockSize(double* parameter_block) const; + int ParameterBlockLocalSize(double* parameter_block) const; + void GetParameterBlocks(vector* parameter_blocks) const; + const Program& program() const { return *program_; } Program* mutable_program() { return program_.get(); } diff --git a/internal/ceres/problem_test.cc b/internal/ceres/problem_test.cc index ab40e05da..0944d3f99 100644 --- a/internal/ceres/problem_test.cc +++ b/internal/ceres/problem_test.cc @@ -502,6 +502,35 @@ TEST(Problem, RemoveParameterBlockWithUnknownPtrDies) { problem.RemoveParameterBlock(y), "Parameter block not found:"); } +TEST(Problem, ParameterBlockQueryTest) { + double x[3]; + double y[4]; + Problem problem; + problem.AddParameterBlock(x, 3); + problem.AddParameterBlock(y, 4); + + vector constant_parameters; + constant_parameters.push_back(0); + problem.SetParameterization( + x, + new SubsetParameterization(3, constant_parameters)); + EXPECT_EQ(problem.ParameterBlockSize(x), 3); + EXPECT_EQ(problem.ParameterBlockLocalSize(x), 2); + EXPECT_EQ(problem.ParameterBlockLocalSize(y), 4); + + vector parameter_blocks; + problem.GetParameterBlocks(¶meter_blocks); + EXPECT_EQ(parameter_blocks.size(), 2); + EXPECT_NE(parameter_blocks[0], parameter_blocks[1]); + EXPECT_TRUE(parameter_blocks[0] == x || parameter_blocks[0] == y); + EXPECT_TRUE(parameter_blocks[1] == x || parameter_blocks[1] == y); + + problem.RemoveParameterBlock(x); + problem.GetParameterBlocks(¶meter_blocks); + EXPECT_EQ(parameter_blocks.size(), 1); + EXPECT_TRUE(parameter_blocks[0] == y); +} + TEST_P(DynamicProblem, RemoveParameterBlockWithNoResiduals) { problem->AddParameterBlock(y, 4); problem->AddParameterBlock(z, 5);