-
Notifications
You must be signed in to change notification settings - Fork 41
Expand file tree
/
Copy pathfactory_resolver.cc
More file actions
227 lines (195 loc) · 8.88 KB
/
Copy pathfactory_resolver.cc
File metadata and controls
227 lines (195 loc) · 8.88 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
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
#include "factory_resolver.h"
#include "api_status.h"
#include "constants.h"
#include "err_constants.h"
#include "logger/event_logger.h"
#include "model_mgmt/empty_data_transport.h"
#include "utility/watchdog.h"
#include "vw_model/pdf_model.h"
#include "vw_model/vw_model.h"
#ifdef RL_BUILD_FEDERATION
# include "federation/federated_loop_controller.h"
#endif
#ifdef USE_AZURE_FACTORIES
# include "azure_factories.h"
# include "model_mgmt/restapi_data_transport.h"
#endif
#include "console_tracer.h"
#include "error_callback_fn.h"
#include "logger/file/file_logger.h"
#include "model_mgmt/file_model_loader.h"
#include <type_traits>
namespace reinforcement_learning
{
namespace m = model_management;
namespace u = utility;
// For proper static intialization
// Check https://en.wikibooks.org/wiki/More_C++_Idioms/Nifty_Counter for explanation
static int init_guard; // guaranteed to be zero when loaded
// properly aligned memory for the factory object
template <typename T>
using natural_align = std::aligned_storage<sizeof(T), alignof(T)>;
static natural_align<data_transport_factory_t>::type dtfactory_buf;
static natural_align<model_factory_t>::type modelfactory_buf;
static natural_align<sender_factory_t>::type senderfactory_buf;
static natural_align<trace_logger_factory_t>::type traceloggerfactory_buf;
static natural_align<time_provider_factory_t>::type time_provider_factory_buf;
// Reference should point to the allocated memory to be initialized by placement new in
// factory_initializer::factory_initializer()
data_transport_factory_t& data_transport_factory = (data_transport_factory_t&)(dtfactory_buf);
model_factory_t& model_factory = (model_factory_t&)(modelfactory_buf);
sender_factory_t& sender_factory = (sender_factory_t&)(senderfactory_buf);
trace_logger_factory_t& trace_logger_factory = (trace_logger_factory_t&)(traceloggerfactory_buf);
time_provider_factory_t& time_provider_factory = (time_provider_factory_t&)(time_provider_factory_buf);
factory_initializer::factory_initializer()
{
if (init_guard++ == 0)
{
new (&data_transport_factory) data_transport_factory_t();
new (&model_factory) model_factory_t();
new (&sender_factory) sender_factory_t();
new (&trace_logger_factory) trace_logger_factory_t();
new (&time_provider_factory) time_provider_factory_t();
register_default_factories();
}
}
factory_initializer::~factory_initializer()
{
if (--init_guard == 0)
{
(&data_transport_factory)->~data_transport_factory_t();
(&model_factory)->~model_factory_t();
(&sender_factory)->~sender_factory_t();
(&trace_logger_factory)->~trace_logger_factory_t();
(&time_provider_factory)->~time_provider_factory_t();
}
}
template <typename model_t>
int model_create(
std::unique_ptr<m::i_model>& retval, const u::configuration& c, i_trace* trace_logger, api_status* status)
{
retval.reset(new model_t(trace_logger, c));
return error_code::success;
}
int null_tracer_create(
std::unique_ptr<i_trace>& retval, const u::configuration& /*cfg*/, i_trace* trace_logger, api_status* status);
int console_tracer_create(
std::unique_ptr<i_trace>& retval, const u::configuration& /*cfg*/, i_trace* trace_logger, api_status* status);
int file_sender_create(std::unique_ptr<i_sender>& retval, const u::configuration& cfg, const char* file_name,
error_callback_fn* error_cb, i_trace* trace_logger, api_status* status)
{
retval.reset(new logger::file::file_logger(file_name, trace_logger));
return error_code::success;
}
int empty_data_transport_create(std::unique_ptr<m::i_data_transport>& retval, const u::configuration& config,
i_trace* trace_logger, api_status* status)
{
TRACE_INFO(trace_logger, "Empty data transport created.");
retval.reset(new model_management::empty_data_transport());
return error_code::success;
}
int file_model_loader_create(std::unique_ptr<m::i_data_transport>& retval, const u::configuration& config,
i_trace* trace_logger, api_status* status)
{
TRACE_INFO(trace_logger, "File model loader created.");
const char* file_name = config.get(name::MODEL_FILE_NAME, "current");
const bool file_must_exist = config.get_bool(name::MODEL_FILE_MUST_EXIST, false);
auto file_loader = VW::make_unique<model_management::file_model_loader>(file_name, file_must_exist, trace_logger);
const auto success = file_loader->init(status);
if (success != error_code::success) { return success; }
retval = std::move(file_loader);
return error_code::success;
}
#ifdef RL_BUILD_FEDERATION
int federated_loop_controller_create(std::unique_ptr<m::i_data_transport>& retval, const u::configuration& config,
i_trace* trace_logger, api_status* status)
{
TRACE_INFO(trace_logger, "Local loop controller i_data_transport created.");
std::unique_ptr<federated_loop_controller> output;
std::unique_ptr<model_management::i_data_transport> transport;
RETURN_IF_FAIL(federated_loop_controller::create(output, config, std::move(transport), trace_logger, status));
retval = std::move(output);
return error_code::success;
}
#endif
int null_time_provider_create(
std::unique_ptr<i_time_provider>& retval, const u::configuration& config, i_trace* trace_logger, api_status* status)
{
TRACE_INFO(trace_logger, "Null time provider created.");
retval.reset();
return error_code::success;
}
int clock_time_provider_create(
std::unique_ptr<i_time_provider>& retval, const u::configuration& config, i_trace* trace_logger, api_status* status)
{
TRACE_INFO(trace_logger, "Clock time provider created.");
retval.reset(new clock_time_provider());
return error_code::success;
}
void factory_initializer::register_default_factories()
{
#ifdef USE_AZURE_FACTORIES
register_azure_factories();
#endif
data_transport_factory.register_type(value::NO_MODEL_DATA, empty_data_transport_create);
data_transport_factory.register_type(value::FILE_MODEL_DATA, file_model_loader_create);
#ifdef RL_BUILD_FEDERATION
data_transport_factory.register_type(value::LOCAL_LOOP_MODEL_DATA, federated_loop_controller_create);
#else
data_transport_factory.register_type(value::LOCAL_LOOP_MODEL_DATA,
[](std::unique_ptr<m::i_data_transport>&, const u::configuration&, i_trace* trace_logger, api_status* status)
{
RETURN_ERROR_ARG(trace_logger, status, create_fn_exception,
"Cannot use LOCAL_LOOP_MODEL_DATA because rlclientlib was not compiled with federated learning enabled");
});
#endif
model_factory.register_type(value::VW, model_create<m::vw_model>);
model_factory.register_type(value::PASSTHROUGH_PDF_MODEL, model_create<m::pdf_model>);
trace_logger_factory.register_type(value::NULL_TRACE_LOGGER, null_tracer_create);
trace_logger_factory.register_type(value::CONSOLE_TRACE_LOGGER, console_tracer_create);
time_provider_factory.register_type(value::NULL_TIME_PROVIDER, null_time_provider_create);
time_provider_factory.register_type(value::CLOCK_TIME_PROVIDER, clock_time_provider_create);
// Register File loggers
sender_factory.register_type(value::EPISODE_FILE_SENDER,
[](std::unique_ptr<i_sender>& retval, const u::configuration& c, error_callback_fn* cb, i_trace* trace_logger,
api_status* status)
{
const char* file_name = c.get(name::EPISODE_FILE_NAME, "episode.fb.data");
return file_sender_create(retval, c, file_name, cb, trace_logger, status);
});
sender_factory.register_type(value::OBSERVATION_FILE_SENDER,
[](std::unique_ptr<i_sender>& retval, const u::configuration& c, error_callback_fn* cb, i_trace* trace_logger,
api_status* status)
{
const char* file_name = c.get(name::OBSERVATION_FILE_NAME, "observation.fb.data");
return file_sender_create(retval, c, file_name, cb, trace_logger, status);
});
sender_factory.register_type(value::INTERACTION_FILE_SENDER,
[](std::unique_ptr<i_sender>& retval, const u::configuration& c, error_callback_fn* cb, i_trace* trace_logger,
api_status* status)
{
const char* file_name = c.get(name::INTERACTION_FILE_NAME, "interaction.fb.data");
return file_sender_create(retval, c, file_name, cb, trace_logger, status);
});
// Register a default factory for LOCAL_LOOP_SENDER that returns an error
sender_factory.register_type(value::LOCAL_LOOP_SENDER,
[](std::unique_ptr<i_sender>&, const u::configuration&, error_callback_fn*, i_trace* trace_logger,
api_status* status)
{
RETURN_ERROR_ARG(trace_logger, status, create_fn_exception,
"LOCAL_LOOP_SENDER must be used with model source set to LOCAL_LOOP_MODEL_DATA");
});
}
int null_tracer_create(
std::unique_ptr<i_trace>& retval, const u::configuration& cfg, i_trace* trace_logger, api_status* status)
{
retval.reset();
return error_code::success;
}
int console_tracer_create(
std::unique_ptr<i_trace>& retval, const u::configuration& cfg, i_trace* trace_logger, api_status* status)
{
retval.reset(new console_tracer());
return error_code::success;
}
} // namespace reinforcement_learning