Skip to content

Navigation Menu

Sign in
Appearance settings

Search code, repositories, users, issues, pull requests...

Provide feedback

We read every piece of feedback, and take your input very seriously.

Saved searches

Use saved searches to filter your results more quickly

Appearance settings

Latest commit

 

History

History
History
104 lines (76 loc) · 2.25 KB

File metadata and controls

104 lines (76 loc) · 2.25 KB
Copy raw file
Download raw file
Open symbols panel
Edit and raw actions
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
#include <cstring>
#include <fstream>
#include <iostream>
#include <string>
#include <iomanip>
#include <memory>
#include <cmath>
#include <stdexcept>
#include <vector>
#include <cstdlib>
#include "fm.h"
using namespace std;
using namespace fm;
struct Option {
string test_path, model_path, output_path;
};
string predict_help() {
return string(
"usage: fm-predict test_file model_file output_file\n");
}
Option parse_option(int argc, char **argv) {
vector<string> args;
for (int i = 0; i < argc; i++)
args.push_back(string(argv[i]));
if (argc == 1)
throw invalid_argument(predict_help());
Option option;
if (argc != 4)
throw invalid_argument("cannot parse argument");
option.test_path = string(args[1]);
option.model_path = string(args[2]);
option.output_path = string(args[3]);
return option;
}
void predict(string test_path, string model_path, string output_path) {
int const kMaxLineSize = 1000000;
FILE *f_in = fopen(test_path.c_str(), "r");
ofstream f_out(output_path);
char line[kMaxLineSize];
fm_model model = fm_load_model(model_path);
fm_double loss = 0;
vector<fm_node> x;
fm_int i = 0;
for (; fgets(line, kMaxLineSize, f_in) != nullptr; i++) {
x.clear();
char *y_char = strtok(line, " \t");
fm_float y = (atoi(y_char)>0)? 1.0f : -1.0f;
while (true) {
char *idx_char = strtok(nullptr,":");
char *value_char = strtok(nullptr," \t");
if (idx_char == nullptr || *idx_char == '\n')
break;
fm_node N;
N.idx = atoi(idx_char);
N.value = atof(value_char);
x.push_back(N);
}
fm_float y_bar = fm_predict(x.data(), x.data()+x.size(), model);
loss -= y==1? log(y_bar) : log(1-y_bar);
f_out << y_bar << "\n";
}
loss /= i;
cout << "logloss = " << fixed << setprecision(5) << loss << endl;
fclose(f_in);
}
int main(int argc, char **argv) {
Option option;
try {
option = parse_option(argc, argv);
} catch(invalid_argument const &e) {
cout << e.what() << endl;
return 1;
}
predict(option.test_path, option.model_path, option.output_path);
return 0;
}
Morty Proxy This is a proxified and sanitized view of the page, visit original site.