autoforge-engine 0.1.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.
modelforge/config.py ADDED
@@ -0,0 +1,580 @@
1
+ from pathlib import Path
2
+ from typing import Any
3
+
4
+ import yaml
5
+
6
+
7
+ class ModelForgeConfig:
8
+ """
9
+ Central configuration manager for ModelForge.
10
+
11
+ Supports:
12
+ - Python dictionary configuration
13
+ - YAML configuration files
14
+ - Configuration validation
15
+ - Configuration persistence
16
+ - Nested feature-selection settings
17
+ - Optimization settings
18
+ - Experiment tracking settings
19
+ """
20
+
21
+ DEFAULTS = {
22
+ "target": None,
23
+ "task_type": None,
24
+ "objective": "balanced",
25
+ "test_size": 0.2,
26
+ "cv": 5,
27
+ "random_state": 42,
28
+ "models": None,
29
+ "excluded_columns": [],
30
+ "feature_selection": {
31
+ "variance_threshold": None,
32
+ "correlation_threshold": None,
33
+ },
34
+ "optimization": {
35
+ "enabled": False,
36
+ "models": 3,
37
+ "max_trials": 10,
38
+ },
39
+ "experiment_directory": ".modelforge/experiments",
40
+ }
41
+
42
+ VALID_TASK_TYPES = {
43
+ "regression",
44
+ "classification",
45
+ }
46
+
47
+ VALID_OBJECTIVES = {
48
+ "balanced",
49
+ "performance",
50
+ "error",
51
+ "speed",
52
+ }
53
+
54
+ def __init__(
55
+ self,
56
+ values: dict[str, Any] | None = None,
57
+ ):
58
+ if values is None:
59
+ values = {}
60
+
61
+ if not isinstance(values, dict):
62
+ raise TypeError(
63
+ "Configuration values must be a dictionary."
64
+ )
65
+
66
+ self.values = self._build_values(values)
67
+ self.validate()
68
+
69
+ @classmethod
70
+ def from_file(
71
+ cls,
72
+ path: str,
73
+ ) -> "ModelForgeConfig":
74
+ """
75
+ Load ModelForge configuration from YAML.
76
+ """
77
+
78
+ config_path = Path(path)
79
+
80
+ if not config_path.exists():
81
+ raise FileNotFoundError(
82
+ f"Configuration file not found: {config_path}"
83
+ )
84
+
85
+ if not config_path.is_file():
86
+ raise ValueError(
87
+ f"Configuration path is not a file: {config_path}"
88
+ )
89
+
90
+ if config_path.suffix.lower() not in {
91
+ ".yaml",
92
+ ".yml",
93
+ }:
94
+ raise ValueError(
95
+ "Configuration file must use "
96
+ ".yaml or .yml extension."
97
+ )
98
+
99
+ with config_path.open(
100
+ "r",
101
+ encoding="utf-8",
102
+ ) as file:
103
+ values = yaml.safe_load(file)
104
+
105
+ if values is None:
106
+ values = {}
107
+
108
+ if not isinstance(values, dict):
109
+ raise ValueError(
110
+ "Configuration root must be a dictionary."
111
+ )
112
+
113
+ return cls(values)
114
+
115
+ def save(
116
+ self,
117
+ path: str,
118
+ overwrite: bool = False,
119
+ ) -> str:
120
+ """
121
+ Save the configuration to YAML.
122
+ """
123
+
124
+ config_path = Path(path)
125
+
126
+ if config_path.exists() and not overwrite:
127
+ raise FileExistsError(
128
+ f"Configuration already exists: {config_path}"
129
+ )
130
+
131
+ config_path.parent.mkdir(
132
+ parents=True,
133
+ exist_ok=True,
134
+ )
135
+
136
+ with config_path.open(
137
+ "w",
138
+ encoding="utf-8",
139
+ ) as file:
140
+ yaml.safe_dump(
141
+ self.values,
142
+ file,
143
+ sort_keys=False,
144
+ )
145
+
146
+ return str(config_path.resolve())
147
+
148
+ def get(
149
+ self,
150
+ key: str,
151
+ default: Any = None,
152
+ ) -> Any:
153
+ """
154
+ Retrieve a top-level configuration value.
155
+ """
156
+
157
+ return self.values.get(
158
+ key,
159
+ default,
160
+ )
161
+
162
+ def get_nested(
163
+ self,
164
+ section: str,
165
+ key: str,
166
+ default: Any = None,
167
+ ) -> Any:
168
+ """
169
+ Retrieve a value from a nested configuration section.
170
+ """
171
+
172
+ section_value = self.values.get(
173
+ section,
174
+ {},
175
+ )
176
+
177
+ if not isinstance(section_value, dict):
178
+ return default
179
+
180
+ return section_value.get(
181
+ key,
182
+ default,
183
+ )
184
+
185
+ def to_dict(self) -> dict[str, Any]:
186
+ """
187
+ Return a deep copy of the configuration.
188
+ """
189
+
190
+ return self._deep_copy(
191
+ self.values
192
+ )
193
+
194
+ def validate(self) -> None:
195
+ """
196
+ Validate all configuration values.
197
+ """
198
+
199
+ target = self.values.get(
200
+ "target"
201
+ )
202
+
203
+ if target is not None:
204
+ if not isinstance(
205
+ target,
206
+ str,
207
+ ):
208
+ raise TypeError(
209
+ "target must be a string."
210
+ )
211
+
212
+ if not target.strip():
213
+ raise ValueError(
214
+ "target cannot be empty."
215
+ )
216
+
217
+ task_type = self.values.get(
218
+ "task_type"
219
+ )
220
+
221
+ if task_type is not None:
222
+ if not isinstance(
223
+ task_type,
224
+ str,
225
+ ):
226
+ raise TypeError(
227
+ "task_type must be a string."
228
+ )
229
+
230
+ if task_type not in self.VALID_TASK_TYPES:
231
+ raise ValueError(
232
+ "task_type must be "
233
+ "'regression' or "
234
+ "'classification'."
235
+ )
236
+
237
+ objective = self.values.get(
238
+ "objective"
239
+ )
240
+
241
+ if objective not in self.VALID_OBJECTIVES:
242
+ raise ValueError(
243
+ "objective must be one of: "
244
+ "balanced, performance, "
245
+ "error, speed."
246
+ )
247
+
248
+ test_size = self.values.get(
249
+ "test_size"
250
+ )
251
+
252
+ if isinstance(test_size, bool):
253
+ raise TypeError(
254
+ "test_size must be numeric."
255
+ )
256
+
257
+ if not isinstance(
258
+ test_size,
259
+ (int, float),
260
+ ):
261
+ raise TypeError(
262
+ "test_size must be numeric."
263
+ )
264
+
265
+ if not 0 < test_size < 1:
266
+ raise ValueError(
267
+ "test_size must be between 0 and 1."
268
+ )
269
+
270
+ cv = self.values.get(
271
+ "cv"
272
+ )
273
+
274
+ if isinstance(cv, bool):
275
+ raise TypeError(
276
+ "cv must be an integer."
277
+ )
278
+
279
+ if not isinstance(
280
+ cv,
281
+ int,
282
+ ):
283
+ raise TypeError(
284
+ "cv must be an integer."
285
+ )
286
+
287
+ if cv < 2:
288
+ raise ValueError(
289
+ "cv must be at least 2."
290
+ )
291
+
292
+ random_state = self.values.get(
293
+ "random_state"
294
+ )
295
+
296
+ if isinstance(random_state, bool):
297
+ raise TypeError(
298
+ "random_state must be an integer."
299
+ )
300
+
301
+ if not isinstance(
302
+ random_state,
303
+ int,
304
+ ):
305
+ raise TypeError(
306
+ "random_state must be an integer."
307
+ )
308
+
309
+ models = self.values.get(
310
+ "models"
311
+ )
312
+
313
+ if models is not None:
314
+ if not isinstance(
315
+ models,
316
+ list,
317
+ ):
318
+ raise TypeError(
319
+ "models must be a list."
320
+ )
321
+
322
+ if not models:
323
+ raise ValueError(
324
+ "models cannot be empty."
325
+ )
326
+
327
+ if not all(
328
+ isinstance(model, str)
329
+ for model in models
330
+ ):
331
+ raise TypeError(
332
+ "Every model name must be a string."
333
+ )
334
+
335
+ excluded_columns = self.values.get(
336
+ "excluded_columns"
337
+ )
338
+
339
+ if not isinstance(
340
+ excluded_columns,
341
+ list,
342
+ ):
343
+ raise TypeError(
344
+ "excluded_columns must be a list."
345
+ )
346
+
347
+ if not all(
348
+ isinstance(column, str)
349
+ for column in excluded_columns
350
+ ):
351
+ raise TypeError(
352
+ "Every excluded column must be a string."
353
+ )
354
+
355
+ feature_selection = self.values.get(
356
+ "feature_selection"
357
+ )
358
+
359
+ if not isinstance(
360
+ feature_selection,
361
+ dict,
362
+ ):
363
+ raise TypeError(
364
+ "feature_selection must be a dictionary."
365
+ )
366
+
367
+ variance_threshold = (
368
+ feature_selection.get(
369
+ "variance_threshold"
370
+ )
371
+ )
372
+
373
+ if variance_threshold is not None:
374
+ if isinstance(
375
+ variance_threshold,
376
+ bool,
377
+ ):
378
+ raise TypeError(
379
+ "variance_threshold must be numeric."
380
+ )
381
+
382
+ if not isinstance(
383
+ variance_threshold,
384
+ (int, float),
385
+ ):
386
+ raise TypeError(
387
+ "variance_threshold must be numeric."
388
+ )
389
+
390
+ if variance_threshold < 0:
391
+ raise ValueError(
392
+ "variance_threshold cannot be negative."
393
+ )
394
+
395
+ correlation_threshold = (
396
+ feature_selection.get(
397
+ "correlation_threshold"
398
+ )
399
+ )
400
+
401
+ if correlation_threshold is not None:
402
+ if isinstance(
403
+ correlation_threshold,
404
+ bool,
405
+ ):
406
+ raise TypeError(
407
+ "correlation_threshold must be numeric."
408
+ )
409
+
410
+ if not isinstance(
411
+ correlation_threshold,
412
+ (int, float),
413
+ ):
414
+ raise TypeError(
415
+ "correlation_threshold must be numeric."
416
+ )
417
+
418
+ if not (
419
+ 0 < correlation_threshold <= 1
420
+ ):
421
+ raise ValueError(
422
+ "correlation_threshold must "
423
+ "be between 0 and 1."
424
+ )
425
+
426
+ optimization = self.values.get(
427
+ "optimization"
428
+ )
429
+
430
+ if not isinstance(
431
+ optimization,
432
+ dict,
433
+ ):
434
+ raise TypeError(
435
+ "optimization must be a dictionary."
436
+ )
437
+
438
+ enabled = optimization.get(
439
+ "enabled"
440
+ )
441
+
442
+ if not isinstance(
443
+ enabled,
444
+ bool,
445
+ ):
446
+ raise TypeError(
447
+ "optimization.enabled must be a boolean."
448
+ )
449
+
450
+ optimization_models = optimization.get(
451
+ "models"
452
+ )
453
+
454
+ if isinstance(
455
+ optimization_models,
456
+ bool,
457
+ ):
458
+ raise TypeError(
459
+ "optimization.models must be an integer."
460
+ )
461
+
462
+ if not isinstance(
463
+ optimization_models,
464
+ int,
465
+ ):
466
+ raise TypeError(
467
+ "optimization.models must be an integer."
468
+ )
469
+
470
+ if optimization_models < 1:
471
+ raise ValueError(
472
+ "optimization.models must be at least 1."
473
+ )
474
+
475
+ max_trials = optimization.get(
476
+ "max_trials"
477
+ )
478
+
479
+ if isinstance(
480
+ max_trials,
481
+ bool,
482
+ ):
483
+ raise TypeError(
484
+ "optimization.max_trials must be an integer."
485
+ )
486
+
487
+ if not isinstance(
488
+ max_trials,
489
+ int,
490
+ ):
491
+ raise TypeError(
492
+ "optimization.max_trials must be an integer."
493
+ )
494
+
495
+ if max_trials < 1:
496
+ raise ValueError(
497
+ "optimization.max_trials must be at least 1."
498
+ )
499
+
500
+ experiment_directory = self.values.get(
501
+ "experiment_directory"
502
+ )
503
+
504
+ if not isinstance(
505
+ experiment_directory,
506
+ str,
507
+ ):
508
+ raise TypeError(
509
+ "experiment_directory must be a string."
510
+ )
511
+
512
+ if not experiment_directory.strip():
513
+ raise ValueError(
514
+ "experiment_directory cannot be empty."
515
+ )
516
+
517
+ @classmethod
518
+ def _build_values(
519
+ cls,
520
+ values: dict[str, Any],
521
+ ) -> dict[str, Any]:
522
+ """
523
+ Merge user configuration with defaults.
524
+ """
525
+
526
+ unknown_keys = (
527
+ set(values) - set(cls.DEFAULTS)
528
+ )
529
+
530
+ if unknown_keys:
531
+ raise ValueError(
532
+ "Unknown configuration keys: "
533
+ + ", ".join(
534
+ sorted(unknown_keys)
535
+ )
536
+ )
537
+
538
+ result = cls._deep_copy(
539
+ cls.DEFAULTS
540
+ )
541
+
542
+ for key, value in values.items():
543
+ if (
544
+ key in {
545
+ "feature_selection",
546
+ "optimization",
547
+ }
548
+ and isinstance(value, dict)
549
+ ):
550
+ result[key].update(value)
551
+ else:
552
+ result[key] = value
553
+
554
+ return result
555
+
556
+ @staticmethod
557
+ def _deep_copy(value):
558
+ if isinstance(
559
+ value,
560
+ dict,
561
+ ):
562
+ return {
563
+ key: ModelForgeConfig._deep_copy(
564
+ item
565
+ )
566
+ for key, item in value.items()
567
+ }
568
+
569
+ if isinstance(
570
+ value,
571
+ list,
572
+ ):
573
+ return [
574
+ ModelForgeConfig._deep_copy(
575
+ item
576
+ )
577
+ for item in value
578
+ ]
579
+
580
+ return value