CLAMS: cleanup eigen_extensions.h, models can be saved/loaded in binary (*.bin) or ascii (*.txt) formats

This commit is contained in:
matlabbe
2016-10-14 15:24:24 -04:00
parent e13454fe3f
commit 0ce6ef8d8d
6 changed files with 129 additions and 223 deletions
@@ -28,25 +28,7 @@ namespace eigen_extensions {
double var = total / (double)vec.rows();
return sqrt(var);
}
template<class S, int T, int U>
void save(const Eigen::Matrix<S, T, U>& mat, const std::string& filename);
template<class S, int T, int U>
void load(const std::string& filename, Eigen::Matrix<S, T, U>* mat);
template<class ScalarType, int Options, class IndexType>
void save(const Eigen::SparseMatrix<ScalarType, Options, IndexType>& mat, const std::string& filename);
template<class ScalarType, int Options, class IndexType>
void load(const std::string& filename, Eigen::SparseMatrix<ScalarType, Options, IndexType>* mat);
template<class S, int T, int U>
void saveASCII(const Eigen::Matrix<S, T, U>& mat, const std::string& filename);
template<class S, int T, int U>
void loadASCII(const std::string& filename, Eigen::Matrix<S, T, U>* mat);
template<class S, int T, int U>
void serialize(const Eigen::Matrix<S, T, U>& mat, std::ostream& strm);
@@ -60,15 +42,6 @@ namespace eigen_extensions {
void deserializeASCII(std::istream& strm, Eigen::Matrix<S, T, U>* mat);
// -- SparseMatrix serialization.
template<class ScalarType, int Options, class IndexType>
void serialize(const Eigen::SparseMatrix<ScalarType, Options, IndexType>& mat, std::ostream& strm);
template<class ScalarType, int Options, class IndexType>
void deserialize(std::istream& strm, Eigen::SparseMatrix<ScalarType, Options, IndexType>* mat);
// -- Scalar serialization
// TODO: Can you name these {de,}serialize() and still have the right
// functions get called when serializing matrices?
@@ -78,6 +51,12 @@ namespace eigen_extensions {
template<class T>
void deserializeScalar(std::istream& strm, T* val);
template<class T>
void serializeScalarASCII(T val, std::ostream& strm);
template<class T>
void deserializeScalarASCII(std::istream& strm, T* val);
/************************************************************
* Template implementations
@@ -111,133 +90,7 @@ namespace eigen_extensions {
*mat = Eigen::Map< Eigen::Matrix<S, T, U> >(buf, rows, cols);
free(buf);
}
/*
template<class S, int T, int U>
void save(const Eigen::Matrix<S, T, U>& mat, const std::string& filename)
{
assert(filename.size() > 3);
if(filename.substr(filename.size() - 3, 3).compare(".gz") == 0) {
ogzstream file(filename.c_str());
assert(file);
serialize(mat, file);
file.close();
}
else {
assert(boost::filesystem::extension(filename).compare(".eig") == 0);
std::ofstream file(filename.c_str());
assert(file);
serialize(mat, file);
file.close();
}
}
template<class S, int T, int U>
void load(const std::string& filename, Eigen::Matrix<S, T, U>* mat)
{
assert(filename.size() > 3);
if(filename.substr(filename.size() - 3, 3).compare(".gz") == 0) {
igzstream file(filename.c_str());
assert(file);
deserialize(file, mat);
file.close();
}
else {
assert(boost::filesystem::extension(filename).compare(".eig") == 0);
std::ifstream file(filename.c_str());
assert(file);
deserialize(file, mat);
file.close();
}
}
*/
template<class ScalarType, int Options, class IndexType>
void serialize(const Eigen::SparseMatrix<ScalarType, Options, IndexType>& mat, std::ostream& strm)
{
int bytes = sizeof(ScalarType);
int type = Options;
int outer = mat.outerSize();
int inner = mat.innerSize();
int nnz = mat.nonZeros();
strm.write((char*)&bytes, sizeof(int));
strm.write((char*)&type, sizeof(int));
strm.write((char*)&outer, sizeof(int));
strm.write((char*)&inner, sizeof(int));
strm.write((char*)&nnz, sizeof(int));
typedef typename Eigen::SparseMatrix<ScalarType, Options, IndexType>::InnerIterator InnerIterator;
for(IndexType i = 0; i < mat.outerSize(); ++i) {
int num = 0;
for(InnerIterator it(mat, i); it; ++it)
++num;
strm.write((const char*)&num, sizeof(num));
for(InnerIterator it(mat, i); it; ++it) {
int idx = it.index();
ScalarType buf = it.value();
strm.write((const char*)&idx, sizeof(idx));
strm.write((const char*)&buf, sizeof(buf));
}
}
}
template<class ScalarType, int Options, class IndexType>
void deserialize(std::istream& strm, Eigen::SparseMatrix<ScalarType, Options, IndexType>* mat)
{
int bytes;
int options;
int outer;
int inner;
int nnz;
strm.read((char*)&bytes, sizeof(int));
strm.read((char*)&options, sizeof(int));
strm.read((char*)&outer, sizeof(int));
strm.read((char*)&inner, sizeof(int));
strm.read((char*)&nnz, sizeof(int));
assert(bytes == sizeof(ScalarType));
assert(options == Options);
if(mat->IsRowMajor)
mat->resize(outer, inner);
else
mat->resize(inner, outer);
mat->reserve(nnz);
ScalarType buf;
for(int i = 0; i < mat->outerSize(); ++i) {
mat->startVec(i);
int num;
strm.read((char*)&num, sizeof(int));
int idx;
for(int j = 0; j < num; ++j) {
strm.read((char*)&idx, sizeof(idx));
strm.read((char*)&buf, sizeof(buf));
mat->insertBackByOuterInner(i, idx) = buf;
}
}
mat->finalize();
}
/*
template<class ScalarType, int Options, class IndexType>
void save(const Eigen::SparseMatrix<ScalarType, Options, IndexType>& mat, const std::string& filename)
{
assert(boost::filesystem::extension(filename).compare(".eig") == 0);
std::ofstream file(filename.c_str());
assert(file);
serialize(mat, file);
file.close();
}
template<class ScalarType, int Options, class IndexType>
void load(const std::string& filename, Eigen::SparseMatrix<ScalarType, Options, IndexType>* mat)
{
assert(filename.size() > 3);
assert(boost::filesystem::extension(filename).compare(".eig") == 0);
std::ifstream file(filename.c_str());
assert(file);
deserialize(file, mat);
file.close();
}
*/
template<class S, int T, int U>
void serializeASCII(const Eigen::Matrix<S, T, U>& mat, std::ostream& strm)
{
@@ -271,30 +124,6 @@ namespace eigen_extensions {
}
}
template<class S, int T, int U>
void saveASCII(const Eigen::Matrix<S, T, U>& mat, const std::string& filename)
{
assert(filename.substr(filename.size() - 8).compare(".eig.txt") == 0);
std::ofstream file;
file.open(filename.c_str());
assert(file);
serializeASCII(mat, file);
file.close();
}
template<class S, int T, int U>
void loadASCII(const std::string& filename, Eigen::Matrix<S, T, U>* mat)
{
assert(filename.substr(filename.size() - 8).compare(".eig.txt") == 0);
std::ifstream file;
file.open(filename.c_str());
if(!file)
std::cerr << "File " << filename << " could not be opened. Dying badly." << std::endl;
assert(file);
deserializeASCII(file, mat);
file.close();
}
template<class T>
void serializeScalar(T val, std::ostream& strm)
{
@@ -306,6 +135,25 @@ namespace eigen_extensions {
{
strm.read((char*)val, sizeof(T));
}
template<class T>
void serializeScalarASCII(T val, std::ostream& strm)
{
int old_precision = strm.precision();
strm.precision(16);
strm << "% " << val << std::endl;
strm.precision(old_precision);
}
template<class T>
void deserializeScalarASCII(std::istream& strm, T* val)
{
std::string line;
while(line.length() == 0) getline(strm, line);
assert(line[0] == '%');
std::istringstream iss(line.substr(1));
iss >> *val;
}
}