-
Notifications
You must be signed in to change notification settings - Fork 388
Expand file tree
/
Copy pathProtoNNModel.cpp
More file actions
143 lines (117 loc) · 4.13 KB
/
Copy pathProtoNNModel.cpp
File metadata and controls
143 lines (117 loc) · 4.13 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT license.
#include "ProtoNN.h"
using namespace EdgeML;
using namespace EdgeML::ProtoNN;
ProtoNNModel::ProtoNNModel(
std::string modelFile)
{
std::ifstream infile(modelFile, std::ios::in|std::ios::binary);
assert(infile.is_open());
// Read the size of the model
size_t modelSize;
infile.read((char*)&modelSize, sizeof(modelSize));
// Allocate buffer
char* buff = new char[modelSize];
//Load model from model file
infile.read((char*)buff, modelSize);
infile.close();
importModel(modelSize, (char *const) buff);
delete[] buff;
}
ProtoNNModel::ProtoNNModel()
{
}
ProtoNNModel::ProtoNNModel(
const size_t numBytes,
const char *const fromModel)
{
importModel(numBytes, fromModel);
}
ProtoNNModel::ProtoNNModel(const int argc, const char** argv)
{
hyperParams.setHyperParamsFromArgs(argc, argv);
params.resizeParamsFromHyperParams(hyperParams);
}
ProtoNNModel::ProtoNNModel(const ProtoNNModel::ProtoNNHyperParams& hyperParams_)
: hyperParams(hyperParams_)
{
LOG_INFO("Initialized model hyperparameters");
params.resizeParamsFromHyperParams(hyperParams_);
LOG_INFO("Resized model parameters");
}
ProtoNNModel::~ProtoNNModel() {}
size_t ProtoNNModel::modelStat()
{
size_t offset = 0;
offset += sizeof(hyperParams);
#ifdef SPARSE_Z_PROTONN
offset += sizeof(bool);
offset += sizeof(Eigen::Index);
Eigen::Index nnz = params.Z.outerIndexPtr()[params.Z.cols()]
- params.Z.outerIndexPtr()[0];
offset += sizeof(FP_TYPE) * nnz;
offset += sizeof(sparseIndex_t) * nnz;
offset += sizeof(sparseIndex_t) * (params.Z.cols() + 1);
#else
offset += sizeof(bool);
offset += sizeof(FP_TYPE) * params.Z.rows() * params.Z.cols();
#endif
offset += sizeof(FP_TYPE) * params.W.rows() * params.W.cols();
offset += sizeof(FP_TYPE) * params.B.rows() * params.B.cols();
return offset;
}
void ProtoNNModel::exportModel(
const size_t modelSize,
char *const toModel)
{
assert(modelSize == modelStat());
size_t offset(0);
#ifdef Z_SPARSE
bool isZSparse(true);
#else
bool isZSparse(false);
#endif
memcpy(toModel + offset, (void *)&hyperParams, sizeof(hyperParams));
offset += sizeof(hyperParams);
#ifdef SPARSE_Z_PROTONN
memcpy(toModel + offset, (void *)&isZSparse, sizeof(bool));
offset += sizeof(bool);
offset += sizeof(Eigen::Index);
Eigen::Index nnz = params.Z.outerIndexPtr()[params.Z.cols()]
- params.Z.outerIndexPtr()[0];
offset += sizeof(FP_TYPE) * nnz;
offset += sizeof(sparseIndex_t) * nnz;
offset += sizeof(sparseIndex_t) * (params.Z.cols() + 1);
///// Resume from here .
#else
memcpy(toModel + offset, (void *)&isZSparse, sizeof(bool));
offset += sizeof(bool);
memcpy(toModel + offset, params.Z.data(), sizeof(FP_TYPE) * params.Z.rows() * params.Z.cols());
offset += sizeof(FP_TYPE) * params.Z.rows() * params.Z.cols();
#endif
memcpy(toModel + offset, params.W.data(), sizeof(FP_TYPE) * params.W.rows() * params.W.cols());
offset += sizeof(FP_TYPE) * params.W.rows() * params.W.cols();
memcpy(toModel + offset, params.B.data(), sizeof(FP_TYPE) * params.B.rows() * params.B.cols());
offset += sizeof(FP_TYPE) * params.B.rows() * params.B.cols();
}
void ProtoNNModel::importModel(const size_t numBytes, const char *const fromModel)
{
size_t offset = 0;
memcpy((void *)&hyperParams, fromModel + offset, sizeof(hyperParams));
offset += sizeof(hyperParams);
params.resizeParamsFromHyperParams(hyperParams, false); // No need to set to zero.
bool isZSparse(true);
#ifdef SPARSE_Z_PROTONN
#else
memcpy((void *)&isZSparse, fromModel + offset, sizeof(bool));
offset += sizeof(bool);
assert(isZSparse == false);
memcpy(params.Z.data(), fromModel + offset, sizeof(FP_TYPE) * params.Z.rows() * params.Z.cols());
offset += sizeof(FP_TYPE) * params.Z.rows() * params.Z.cols();
#endif
memcpy(params.W.data(), fromModel + offset, sizeof(FP_TYPE) * params.W.rows() * params.W.cols());
offset += sizeof(FP_TYPE) * params.W.rows() * params.W.cols();
memcpy(params.B.data(), fromModel + offset, sizeof(FP_TYPE) * params.B.rows() * params.B.cols());
offset += sizeof(FP_TYPE) * params.B.rows() * params.B.cols();
}