experiment_server 0.3.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,348 @@
1
+ from pathlib import Path
2
+ import random
3
+ from shutil import ExecError
4
+ from typing import Any, Callable, Dict, List, Tuple, Union
5
+
6
+ from tornado.locale import load_translations
7
+ from experiment_server._participant_ordering import construct_participant_condition, ORDERING_BEHAVIOUR
8
+ from experiment_server.utils import ExperimentServerConfigurationExcetion, ExperimentServerExcetion, merge_dicts
9
+ from loguru import logger
10
+ from easydict import EasyDict as edict
11
+ import json
12
+ import toml
13
+
14
+
15
+ TOP_LEVEL_RESERVED_KEYS = ["name", "config", "extends"]
16
+ SECTIONS = ["main_configuration", "init_configuration", "final_configuration", "template_values", "order", "settings"]
17
+ ALLOWED_SETTINGS = ["randomize_within_groups", "randomize_groups"]
18
+
19
+
20
+ def get_sections(f: Union[str, Path]) -> Dict[str, str]:
21
+ loaded_configurations = {}
22
+ current_blob = None
23
+ current_section = None
24
+ with open(f) as fp:
25
+ for line in fp.readlines():
26
+ stripped_line = line.strip()
27
+ if stripped_line.startswith("//"):
28
+ stripped_line = stripped_line.lstrip("//")
29
+ if stripped_line in SECTIONS: # Lines starting with // are considered comments
30
+ if current_blob is not None:
31
+ loaded_configurations[current_section] = current_blob
32
+ current_blob = ""
33
+ current_section = stripped_line
34
+ else:
35
+ try:
36
+ line = line.rstrip()
37
+ if len(line) > 0:
38
+ current_blob += line
39
+ except TypeError:
40
+ raise ExperimentServerConfigurationExcetion("The file should start with a section header.")
41
+
42
+ loaded_configurations[current_section] = current_blob
43
+
44
+ logger.info(f"Sections in configuration: {list(loaded_configurations.keys())}")
45
+ return loaded_configurations
46
+
47
+
48
+ def process_config_file(f: Union[str, Path], participant_index: int, supress_message:bool=False) -> List[Dict[str, Any]]:
49
+ if participant_index < 1:
50
+ raise ExperimentServerConfigurationExcetion(f"Participant index needs to be greater than 0, got {participant_index}")
51
+
52
+ if Path(f).suffix == ".expconfig":
53
+ return _process_expconfig(f, participant_index, supress_message)
54
+ elif Path(f).suffix == ".toml":
55
+ return _process_toml(f, participant_index, supress_message)
56
+ else:
57
+ raise ExperimentServerExcetion("Invalid file type. Expected `.expconfig` or `.toml`")
58
+
59
+
60
+ def _process_toml(f: Union[str, Path], participant_index:int, supress_message:bool=False) -> List[Dict[str, Any]]:
61
+ loaded_configuration = toml.load(f)
62
+
63
+ configurations = loaded_configuration.get("configuration", {})
64
+ variables = configurations.get("variables", {})
65
+
66
+ configurations_groups = configurations.get("groups", ORDERING_BEHAVIOUR.as_is)
67
+ configurations_within_groups = configurations.get("within_groups", ORDERING_BEHAVIOUR.as_is)
68
+
69
+ random_seed = configurations.get("random_seed", 0)
70
+ random.seed(random_seed + participant_index)
71
+
72
+ all_blocks = _replace_variables(loaded_configuration["blocks"], variables)
73
+ order = configurations.get("order", [list(range(len(all_blocks)))])
74
+
75
+ for c in all_blocks:
76
+ c["name"] = str(c["name"])
77
+
78
+ blocks = construct_participant_condition(all_blocks, participant_index, order=order,
79
+ groups=configurations_groups,
80
+ within_groups=configurations_within_groups)
81
+
82
+ init_blocks = _replace_variables(loaded_configuration.get("init_blocks", []), variables)
83
+ final_blocks = _replace_variables(loaded_configuration.get("final_blocks", []), variables)
84
+
85
+ block_names = [c["name"] for c in (init_blocks + blocks + final_blocks)]
86
+ # Using merge_dicts to ensure the values are references
87
+ resolved_blocks = {c["name"]: c for c in
88
+ resolve_extends([merge_dicts(b, {}) for b in
89
+ (init_blocks + all_blocks + final_blocks)])}
90
+
91
+ # Use block names to get resolved blocks in the expected order
92
+ blocks = [resolved_blocks[c] for c in block_names]
93
+ blocks = resolve_function_calls(blocks)
94
+
95
+ for (idx, c) in enumerate(blocks):
96
+ c["config"]["participant_index"] = participant_index
97
+ c["config"]["name"] = c["name"]
98
+ c["config"]["block_id"] = idx
99
+
100
+ if not supress_message:
101
+ logger.info("Configuration loaded: \n" + json.dumps(blocks, indent=2))
102
+ return blocks
103
+
104
+
105
+ def _process_expconfig(f: Union[str, Path], participant_index: int, supress_message:bool=False) -> List[Dict[str, Any]]:
106
+ loaded_configurations = get_sections(f)
107
+ if "template_values" in loaded_configurations:
108
+ template_values = json.loads(loaded_configurations["template_values"])
109
+ else:
110
+ template_values = {}
111
+
112
+ # config = json.loads(_replace_template_values(loaded_configurations["init_configuration"], template_values))
113
+ if "order" in loaded_configurations:
114
+ order = json.loads(loaded_configurations["order"])
115
+ else:
116
+ order = [list(range(len(loaded_configurations)))]
117
+
118
+ if "settings" in loaded_configurations:
119
+ settings = edict(json.loads(loaded_configurations["settings"]))
120
+ else:
121
+ settings = edict()
122
+
123
+ settings.groups = settings.get("groups", ORDERING_BEHAVIOUR.as_is)
124
+ settings.within_groups = settings.get("within_groups", ORDERING_BEHAVIOUR.as_is)
125
+
126
+ logger.info(f"Settings used: \n {json.dumps(settings, indent=4)}")
127
+
128
+ _raw_main_configuration = loaded_configurations["main_configuration"]
129
+ _templated_main_configuration = _replace_template_values(_raw_main_configuration, template_values)
130
+ try:
131
+ main_configuration = json.loads(_templated_main_configuration)
132
+ except Exception as e:
133
+ logger.error("Raw main config: " + _raw_main_configuration)
134
+ logger.error("Main config with template values passed: " + _templated_main_configuration)
135
+ if isinstance(e, json.decoder.JSONDecodeError):
136
+ raise ExperimentServerConfigurationExcetion("JSONDecodeError at position {}: `... {} ...`".format(
137
+ e.pos,
138
+ _templated_main_configuration[max(0, e.pos - 40):min(len(_templated_main_configuration), e.pos + 40)]))
139
+ else:
140
+ raise
141
+ main_configuration = construct_participant_condition(main_configuration, participant_index, order=order,
142
+ groups=settings.groups,
143
+ within_groups=settings.within_groups)
144
+
145
+ if "init_configuration" in loaded_configurations:
146
+ init_configuration = json.loads(_replace_template_values(loaded_configurations["init_configuration"], template_values))
147
+ else:
148
+ init_configuration = []
149
+
150
+ if "final_configuration" in loaded_configurations:
151
+ final_configuration = json.loads(_replace_template_values(loaded_configurations["final_configuration"], template_values))
152
+ else:
153
+ final_configuration = []
154
+
155
+ config = init_configuration + main_configuration + final_configuration
156
+ config = resolve_extends(config)
157
+
158
+ try:
159
+ for (idx, c) in enumerate(config):
160
+ c["config"]["participant_index"] = participant_index
161
+ c["config"]["name"] = c["name"]
162
+ c["config"]["block_id"] = idx
163
+ except KeyError:
164
+ raise ExperimentServerConfigurationExcetion("blocks missing keys (config/name)")
165
+
166
+ if not supress_message:
167
+ logger.info("Configuration loaded: \n" + json.dumps(config, indent=2))
168
+ return config
169
+
170
+
171
+ def _resolve_extends(c, configs, seen_configs):
172
+ """
173
+ Recursively go through all dependents and collect all values.
174
+ If cyclic dependancy is encountered, it will simply merge all configs along the dependancy path.
175
+ """
176
+ if "extends" not in c or c["extends"] in seen_configs:
177
+ return c, configs, seen_configs
178
+
179
+ dict_a = c
180
+ try:
181
+ dict_b = [_c for _c in configs if _c["name"] == c["extends"]][0]
182
+ except IndexError:
183
+ raise ExperimentServerConfigurationExcetion("`{}` is not a valid name. It must be a `name`.".format(c["extends"]))
184
+
185
+ dict_b, configs, seen_configs = _resolve_extends(dict_b, configs, seen_configs + [dict_b["name"]])
186
+ configs[dict_a["block_idx"]] = merge_dicts(dict_a, dict_b)
187
+ return configs[dict_a["block_idx"]], configs, seen_configs
188
+
189
+
190
+ def resolve_extends(configs):
191
+ # Adding idx to track the configs
192
+ for idx, c in enumerate(configs):
193
+ c["block_idx"] = idx
194
+
195
+ for c in configs:
196
+ configs[c["block_idx"]] = _resolve_extends(c, configs, [c["name"]])[0]
197
+
198
+ # Removing idx
199
+ for c in configs:
200
+ del c["block_idx"]
201
+
202
+ return configs
203
+
204
+
205
+ def _replace_template_values(string, template_values):
206
+ for k, v in template_values.items():
207
+ string = string.replace("{" + k + "}", json.dumps(v))
208
+ return string
209
+
210
+
211
+ def _replace_variables(config: Union[Dict[str, Any], List[Any]], variabels: Dict[str, Any]) -> Union[Dict[str, Any], List[Any]]:
212
+ resolved_config: Union[Dict, List]
213
+ if isinstance(config, dict):
214
+ resolved_config = {}
215
+ for k, v in config.items():
216
+ if isinstance(v, str) and v.startswith("$"):
217
+ try:
218
+ resolved_config[k] = variabels[v[1:]]
219
+ except KeyError:
220
+ raise ExperimentServerConfigurationExcetion(f"The variable `{v}` does not exsist in `configuration.variables`")
221
+ elif isinstance(v, dict):
222
+ resolved_config[k] = _replace_variables(v, variabels)
223
+ else:
224
+ resolved_config[k] = v
225
+ elif isinstance(config, list):
226
+ resolved_config = []
227
+ for v in config:
228
+ if isinstance(v, str) and v.startswith("$"):
229
+ try:
230
+ resolved_config.append(variabels[v[1:]])
231
+ except KeyError:
232
+ raise ExperimentServerConfigurationExcetion(f"The variable `{v}` does not exsist in `configuration.variables`")
233
+ elif isinstance(v, dict):
234
+ resolved_config.append(_replace_variables(v, variabels))
235
+ else:
236
+ resolved_config.append(v)
237
+ return resolved_config
238
+
239
+
240
+ def resolve_function_calls(configs: list) -> list:
241
+ """Check all function calls and replace the values with the result of the function calls."""
242
+ function_calls: Dict[Any, Any] = {}
243
+ return [_resolve_function_calls(c, function_calls) for c in configs]
244
+
245
+
246
+ def _resolve_function_calls(config: dict, function_calls: dict):
247
+ """Recursive function to go traverse through tree and resolve functions."""
248
+ resolved_config = {}
249
+ for k, v in config.items():
250
+ if isinstance(v, dict):
251
+ if len(v) in (2, 3, 4) and all([_k in ["function_name", "args", "params", "id"] for _k in v.keys()]):
252
+ resolved_config[k] = _resolve_function(**v, function_calls=function_calls)
253
+ else:
254
+ resolved_config[k] = _resolve_function_calls(v, function_calls)
255
+ else:
256
+ resolved_config[k] = v
257
+ return resolved_config
258
+
259
+
260
+ def _unpack_args(args) -> Tuple[list, dict]:
261
+ """Convert args into list or dict to allow unpacking."""
262
+ largs, kwargs = [], {}
263
+ if isinstance(args, list):
264
+ largs = args
265
+ elif isinstance(args, dict):
266
+ kwargs = args
267
+ else:
268
+ raise ExperimentServerConfigurationExcetion(f"`args` should be a list or a dict. Got {args}")
269
+ return largs, kwargs
270
+
271
+
272
+ def _resolve_function(function_name:str, args: Union[List,Dict], function_calls: dict, params: Any=None, id: Any=None) -> Any:
273
+ """Call the function and return the value."""
274
+ if id is None:
275
+ call_signature = hash(json.dumps({"function_name": function_name, "args": args, "params": params}, sort_keys=True))
276
+ else:
277
+ call_signature = id
278
+
279
+ if function_name == "choices":
280
+ try:
281
+ function_call_group = function_calls[call_signature]
282
+ except KeyError:
283
+ function_call_group = function_calls[call_signature] = ChoicesFunction(args, params)
284
+ return function_call_group(args, params)
285
+ else:
286
+ raise ExperimentServerConfigurationExcetion(f"Unknown function {function_name}")
287
+
288
+
289
+ class ChoicesFunction:
290
+ """Wrapper for random.choices function call.
291
+ `args` will be passed to `random.choices`.
292
+ If `params` has `unique` whose value is True, will ensure no duplicate values seen in any of the choices call."""
293
+ def __init__(self, args, params) -> None:
294
+ self.args = args
295
+ self.largs, self.kwargs = _unpack_args(args)
296
+ self.unique = False
297
+ self.params = params
298
+ if params is not None:
299
+ if not isinstance(params, dict):
300
+ raise ExperimentServerConfigurationExcetion(f"`params` for `choices` should be a dict.")
301
+ if len(params) not in (0, 1):
302
+ raise ExperimentServerConfigurationExcetion(f"Function `choices` expected 0 or 1 keys in params, got {len(params)}")
303
+ if len(params) == 1 and "unique" not in params:
304
+ raise ExperimentServerConfigurationExcetion(f"Unexpected key in `params` of `choices`. Allowed keys: [`unique`]")
305
+ if "unique" in params:
306
+ self.unique = params.get("unique")
307
+ self.previous_choices = []
308
+
309
+ def __call__(self, args, params) -> Any:
310
+ # Sanity check, making sure nothing changes between calls
311
+ assert self.args == args
312
+ assert params == self.params
313
+ choice = random.choices(*self.largs, **self.kwargs)
314
+ if self.unique:
315
+ # Making sure there are only unique values
316
+ i = 0
317
+ while any([c in self.previous_choices for c in choice]) or len(set(choice)) != len(choice):
318
+ if i > 20:
319
+ # KLUDGE: Chouldn't find unique values?
320
+ break
321
+ i += 1
322
+ choice = random.choices(*self.largs, **self.kwargs)
323
+
324
+ self.previous_choices.extend(choice)
325
+
326
+ # Check if it is possible to get unique values.
327
+ if len(self.previous_choices) != len(set(self.previous_choices)):
328
+ raise ExperimentServerConfigurationExcetion("There are more calls to `choices` than number of elements in `args`")
329
+ return choice
330
+
331
+
332
+ def verify_config(f: Union[str, Path], test_func:Callable[[List[Dict[str, Any]]], Tuple[bool, str]]=None) -> bool:
333
+ import pandas as pd
334
+ from tabulate import tabulate
335
+ with logger.catch(reraise=False, message="Config verification failed"):
336
+ config_blocks = {}
337
+ for participant_index in range(1,6):
338
+ config = process_config_file(f, participant_index=participant_index)
339
+ config_blocks[participant_index] = {f"trial_{idx + 1}": c["name"] for idx, c in enumerate(config)}
340
+ if test_func is not None:
341
+ test_result, reason = test_func(config)
342
+ assert test_result, f"test_func failed for {participant_index} with reason, {reason}"
343
+ df = pd.DataFrame(config_blocks)
344
+ df.style.set_properties(**{'text-align': 'left'}).set_table_styles([ dict(selector='th', props=[('text-align', 'left')])])
345
+ logger.info(f"Ordering for 5 participants: \n\n{tabulate(df, headers='keys', tablefmt='fancy_grid')}\n")
346
+ logger.info(f"Config file verification successful for {f}")
347
+ return True
348
+ return False