jupytermind 0.3.0

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.
Files changed (193) hide show
  1. package/.github/skills/ai-chemistry-scientist/SKILL.md +97 -0
  2. package/.github/skills/ai-chemistry-scientist/manifest.json +156 -0
  3. package/.github/skills/ai-data-scientist/SKILL.md +330 -0
  4. package/.github/skills/ai-genomics-scientist/SKILL.md +98 -0
  5. package/.github/skills/ai-genomics-scientist/manifest.json +93 -0
  6. package/.github/skills/ai-materials-scientist/SKILL.md +51 -0
  7. package/.github/skills/ai-materials-scientist/manifest.json +58 -0
  8. package/.github/skills/ai-scientist/SKILL.md +69 -0
  9. package/.github/skills/ai-scientist/manifest.json +61 -0
  10. package/.github/skills/ai-structural-biology-scientist/SKILL.md +67 -0
  11. package/.github/skills/ai-structural-biology-scientist/manifest.json +72 -0
  12. package/.github/skills/japanese-prose/NOTICE.md +17 -0
  13. package/.github/skills/japanese-prose/SKILL.md +111 -0
  14. package/.github/skills/japanese-prose/references/review-workflow.md +50 -0
  15. package/.github/skills/japanese-prose/references/scoring.md +24 -0
  16. package/.github/skills/japanese-prose/references/writing-guidelines.md +60 -0
  17. package/.github/skills/japanese-prose/scripts/core.py +192 -0
  18. package/.github/skills/japanese-prose/scripts/fixtures/natural.md +5 -0
  19. package/.github/skills/japanese-prose/scripts/fixtures/unnatural.md +5 -0
  20. package/.github/skills/japanese-prose/scripts/lint.py +378 -0
  21. package/.github/skills/japanese-prose/scripts/outline.py +68 -0
  22. package/.github/skills/japanese-prose/scripts/terms.py +112 -0
  23. package/.github/skills/japanese-prose/scripts/test_engine.py +117 -0
  24. package/.github/skills/presentation-planner/SKILL.md +257 -0
  25. package/.github/skills/presentation-planner/assets/design-templates/data-report.yaml +97 -0
  26. package/.github/skills/presentation-planner/assets/design-templates/executive-proposal.yaml +92 -0
  27. package/.github/skills/presentation-planner/assets/design-templates/technical-briefing.yaml +96 -0
  28. package/.github/skills/presentation-planner/assets/scenario-templates/data-report.md +47 -0
  29. package/.github/skills/presentation-planner/assets/scenario-templates/executive-decision.md +43 -0
  30. package/.github/skills/presentation-planner/assets/scenario-templates/technical-briefing.md +45 -0
  31. package/.github/skills/presentation-planner/references/customizing-design-templates.md +160 -0
  32. package/.github/skills/presentation-planner/references/design-spec-schema.md +72 -0
  33. package/.github/skills/presentation-planner/references/handoff-contract.md +49 -0
  34. package/.github/skills/presentation-planner/references/responsibility-boundary.md +32 -0
  35. package/.github/skills/presentation-planner/references/scenario-templates.md +55 -0
  36. package/.github/skills/tech-writer/SKILL.md +434 -0
  37. package/.github/skills/tech-writer/assets/templates/blueprint.md +187 -0
  38. package/.github/skills/tech-writer/assets/templates/design-doc.md +29 -0
  39. package/.github/skills/tech-writer/assets/templates/migration-plan.md +173 -0
  40. package/.github/skills/tech-writer/assets/templates/operations-runbook.md +202 -0
  41. package/.github/skills/tech-writer/assets/templates/pr-description.md +23 -0
  42. package/.github/skills/tech-writer/assets/templates/qiita.md +44 -0
  43. package/.github/skills/tech-writer/assets/templates/readme.md +38 -0
  44. package/.github/skills/tech-writer/assets/templates/requirements-definition.md +170 -0
  45. package/.github/skills/tech-writer/assets/templates/rfi.md +113 -0
  46. package/.github/skills/tech-writer/assets/templates/rfp.md +180 -0
  47. package/.github/skills/tech-writer/assets/templates/security-design.md +167 -0
  48. package/.github/skills/tech-writer/assets/templates/system-design.md +220 -0
  49. package/.github/skills/tech-writer/assets/templates/technical-proposal.md +112 -0
  50. package/.github/skills/tech-writer/assets/templates/test-plan.md +153 -0
  51. package/.github/skills/tech-writer/assets/templates/user-manual.md +22 -0
  52. package/.github/skills/tech-writer/assets/templates/white-paper.md +192 -0
  53. package/.github/skills/tech-writer/references/doctypes/api-docs.md +33 -0
  54. package/.github/skills/tech-writer/references/doctypes/blueprint.md +81 -0
  55. package/.github/skills/tech-writer/references/doctypes/code-comments.md +39 -0
  56. package/.github/skills/tech-writer/references/doctypes/design-doc.md +42 -0
  57. package/.github/skills/tech-writer/references/doctypes/migration-plan.md +63 -0
  58. package/.github/skills/tech-writer/references/doctypes/operations-runbook.md +63 -0
  59. package/.github/skills/tech-writer/references/doctypes/pr-commit.md +82 -0
  60. package/.github/skills/tech-writer/references/doctypes/qiita.md +75 -0
  61. package/.github/skills/tech-writer/references/doctypes/readme.md +43 -0
  62. package/.github/skills/tech-writer/references/doctypes/release-notes.md +30 -0
  63. package/.github/skills/tech-writer/references/doctypes/requirements-definition.md +61 -0
  64. package/.github/skills/tech-writer/references/doctypes/rfi.md +43 -0
  65. package/.github/skills/tech-writer/references/doctypes/rfp.md +46 -0
  66. package/.github/skills/tech-writer/references/doctypes/security-design.md +71 -0
  67. package/.github/skills/tech-writer/references/doctypes/system-design.md +74 -0
  68. package/.github/skills/tech-writer/references/doctypes/technical-proposal.md +49 -0
  69. package/.github/skills/tech-writer/references/doctypes/test-plan.md +67 -0
  70. package/.github/skills/tech-writer/references/doctypes/user-manual.md +58 -0
  71. package/.github/skills/tech-writer/references/doctypes/white-paper.md +84 -0
  72. package/.github/skills/tech-writer/references/doctypes/zenn.md +66 -0
  73. package/.github/skills/tech-writer/references/japanese-prose-optimization.md +110 -0
  74. package/.github/skills/tech-writer/references/style-constitution.md +104 -0
  75. package/.github/skills/tech-writer/scripts/lint.py +412 -0
  76. package/LICENSE +21 -0
  77. package/README.md +92 -0
  78. package/bin/ai-data-scientist.js +123 -0
  79. package/package.json +41 -0
  80. package/pyproject.toml +45 -0
  81. package/src/ai_chemistry_scientist/__init__.py +0 -0
  82. package/src/ai_chemistry_scientist/admet_prediction.py +71 -0
  83. package/src/ai_chemistry_scientist/bioactivity_classification.py +73 -0
  84. package/src/ai_chemistry_scientist/data/sample_molecules.csv +21 -0
  85. package/src/ai_chemistry_scientist/dispatch.py +369 -0
  86. package/src/ai_chemistry_scientist/docking_score.py +97 -0
  87. package/src/ai_chemistry_scientist/drug_likeness_rules.py +84 -0
  88. package/src/ai_chemistry_scientist/evidence.py +41 -0
  89. package/src/ai_chemistry_scientist/molecular_descriptors.py +97 -0
  90. package/src/ai_chemistry_scientist/molecular_formula_mass.py +40 -0
  91. package/src/ai_chemistry_scientist/molecular_similarity.py +78 -0
  92. package/src/ai_chemistry_scientist/qsar_modeling.py +105 -0
  93. package/src/ai_chemistry_scientist/salt_standardization.py +81 -0
  94. package/src/ai_chemistry_scientist/structural_alerts.py +76 -0
  95. package/src/ai_chemistry_scientist/structure_format_conversion.py +84 -0
  96. package/src/ai_chemistry_scientist/validation.py +70 -0
  97. package/src/ai_data_scientist/__init__.py +0 -0
  98. package/src/ai_data_scientist/analysis_assumptions.py +121 -0
  99. package/src/ai_data_scientist/anomaly_detection.py +39 -0
  100. package/src/ai_data_scientist/automl.py +109 -0
  101. package/src/ai_data_scientist/cleaning.py +56 -0
  102. package/src/ai_data_scientist/cli.py +90 -0
  103. package/src/ai_data_scientist/clustering.py +54 -0
  104. package/src/ai_data_scientist/dashboard.py +33 -0
  105. package/src/ai_data_scientist/data_definition.py +100 -0
  106. package/src/ai_data_scientist/data_quality.py +164 -0
  107. package/src/ai_data_scientist/dataset_validation.py +135 -0
  108. package/src/ai_data_scientist/dependency_pins.py +60 -0
  109. package/src/ai_data_scientist/eda.py +82 -0
  110. package/src/ai_data_scientist/experiment_evaluation.py +635 -0
  111. package/src/ai_data_scientist/explainability.py +340 -0
  112. package/src/ai_data_scientist/feature_engineering.py +163 -0
  113. package/src/ai_data_scientist/gate_config.py +32 -0
  114. package/src/ai_data_scientist/ingestion.py +127 -0
  115. package/src/ai_data_scientist/insight_engine.py +180 -0
  116. package/src/ai_data_scientist/japanese_nlp.py +43 -0
  117. package/src/ai_data_scientist/jupyter_launcher.py +137 -0
  118. package/src/ai_data_scientist/jupyter_mcp_client.py +94 -0
  119. package/src/ai_data_scientist/language_router.py +28 -0
  120. package/src/ai_data_scientist/lifecycle.py +221 -0
  121. package/src/ai_data_scientist/mcp_gateway.py +113 -0
  122. package/src/ai_data_scientist/mcp_runtime.py +194 -0
  123. package/src/ai_data_scientist/mcp_transport.py +53 -0
  124. package/src/ai_data_scientist/ml_modeling.py +451 -0
  125. package/src/ai_data_scientist/model_tuning.py +104 -0
  126. package/src/ai_data_scientist/notebook_audit.py +574 -0
  127. package/src/ai_data_scientist/project_manager.py +243 -0
  128. package/src/ai_data_scientist/report_export.py +73 -0
  129. package/src/ai_data_scientist/sensitivity.py +445 -0
  130. package/src/ai_data_scientist/signal_analysis.py +201 -0
  131. package/src/ai_data_scientist/skill_packaging.py +40 -0
  132. package/src/ai_data_scientist/stats_analysis.py +88 -0
  133. package/src/ai_data_scientist/text_nlp.py +44 -0
  134. package/src/ai_data_scientist/timeseries.py +68 -0
  135. package/src/ai_data_scientist/visualization.py +708 -0
  136. package/src/ai_genomics_scientist/__init__.py +1 -0
  137. package/src/ai_genomics_scientist/differential_expression.py +147 -0
  138. package/src/ai_genomics_scientist/dispatch.py +267 -0
  139. package/src/ai_genomics_scientist/evidence.py +45 -0
  140. package/src/ai_genomics_scientist/gene_set_enrichment.py +76 -0
  141. package/src/ai_genomics_scientist/sequence_alignment.py +97 -0
  142. package/src/ai_genomics_scientist/sequence_features.py +111 -0
  143. package/src/ai_genomics_scientist/splice_site_scoring.py +66 -0
  144. package/src/ai_genomics_scientist/validation.py +83 -0
  145. package/src/ai_genomics_scientist/variant_effect.py +147 -0
  146. package/src/ai_genomics_scientist/variant_pathogenicity.py +125 -0
  147. package/src/ai_materials_scientist/__init__.py +0 -0
  148. package/src/ai_materials_scientist/calphad.py +117 -0
  149. package/src/ai_materials_scientist/classical_monte_carlo.py +165 -0
  150. package/src/ai_materials_scientist/crystal_plasticity.py +184 -0
  151. package/src/ai_materials_scientist/dispatch.py +100 -0
  152. package/src/ai_materials_scientist/evidence.py +84 -0
  153. package/src/ai_materials_scientist/fem.py +279 -0
  154. package/src/ai_materials_scientist/kinetic_monte_carlo.py +145 -0
  155. package/src/ai_materials_scientist/molecular_dynamics.py +240 -0
  156. package/src/ai_materials_scientist/phase_field.py +167 -0
  157. package/src/ai_materials_scientist/validation.py +70 -0
  158. package/src/ai_scientist/__init__.py +1 -0
  159. package/src/ai_scientist/completion_gate.py +15 -0
  160. package/src/ai_scientist/data_analysis.py +46 -0
  161. package/src/ai_scientist/evidence_registry.py +99 -0
  162. package/src/ai_scientist/experimental_design.py +20 -0
  163. package/src/ai_scientist/language.py +14 -0
  164. package/src/ai_scientist/latex_renderer.py +41 -0
  165. package/src/ai_scientist/literature_review.py +37 -0
  166. package/src/ai_scientist/manifest.py +87 -0
  167. package/src/ai_scientist/manuscript.py +94 -0
  168. package/src/ai_scientist/mcp_config.py +76 -0
  169. package/src/ai_scientist/mcp_external.py +42 -0
  170. package/src/ai_scientist/mcp_failures.py +23 -0
  171. package/src/ai_scientist/mcp_gateway.py +38 -0
  172. package/src/ai_scientist/mcp_managed.py +180 -0
  173. package/src/ai_scientist/npm_packaging.py +49 -0
  174. package/src/ai_scientist/orchestrator.py +133 -0
  175. package/src/ai_scientist/peer_review.py +60 -0
  176. package/src/ai_scientist/phase_gate.py +74 -0
  177. package/src/ai_scientist/phase_state.py +230 -0
  178. package/src/ai_scientist/presentation.py +56 -0
  179. package/src/ai_scientist/project_config.py +31 -0
  180. package/src/ai_scientist/project_handle.py +74 -0
  181. package/src/ai_scientist/reproducibility.py +20 -0
  182. package/src/ai_scientist/research_planning.py +20 -0
  183. package/src/ai_scientist/skill_invocation.py +21 -0
  184. package/src/ai_scientist/tdd_gate.py +99 -0
  185. package/src/ai_structural_biology_scientist/__init__.py +0 -0
  186. package/src/ai_structural_biology_scientist/contact_map.py +87 -0
  187. package/src/ai_structural_biology_scientist/dispatch.py +269 -0
  188. package/src/ai_structural_biology_scientist/evidence.py +43 -0
  189. package/src/ai_structural_biology_scientist/hydrophobicity.py +101 -0
  190. package/src/ai_structural_biology_scientist/protein_docking_score.py +104 -0
  191. package/src/ai_structural_biology_scientist/secondary_structure.py +95 -0
  192. package/src/ai_structural_biology_scientist/structural_similarity.py +74 -0
  193. package/src/ai_structural_biology_scientist/validation.py +100 -0
@@ -0,0 +1,451 @@
1
+ """Supervised ML modeling.
2
+
3
+ Implements DES-AIDS-012 (REQ-AIDS-008): splits a dataframe into train/test
4
+ partitions, trains the requested classification or regression model, and
5
+ reports the appropriate evaluation metrics for that model type.
6
+
7
+ CHANGE-009 (REQ-AIDS-074..078) extends this module with pluggable
8
+ cv_strategy/cv_splits/scoring and estimator-resolution support; see
9
+ CODE-AIDS-119..122 for the annotated extension points.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ from collections import Counter
15
+ from collections.abc import Callable, Sequence
16
+ from dataclasses import dataclass
17
+ from typing import Any
18
+
19
+ import numpy as np
20
+ import pandas as pd
21
+ from sklearn.base import clone
22
+ from sklearn.ensemble import (
23
+ GradientBoostingClassifier,
24
+ GradientBoostingRegressor,
25
+ RandomForestClassifier,
26
+ RandomForestRegressor,
27
+ )
28
+ from sklearn.linear_model import LinearRegression, LogisticRegression
29
+ from sklearn.metrics import (
30
+ accuracy_score,
31
+ log_loss,
32
+ mean_squared_error,
33
+ precision_score,
34
+ r2_score,
35
+ recall_score,
36
+ roc_auc_score,
37
+ )
38
+ from sklearn.model_selection import GroupKFold, KFold, StratifiedKFold, train_test_split
39
+
40
+ MODEL_BUILDERS = {
41
+ "classification": {
42
+ "random_forest": lambda **params: RandomForestClassifier(random_state=42, **params),
43
+ "logistic_regression": lambda **params: LogisticRegression(max_iter=1000, **params),
44
+ "gradient_boosting": lambda **params: GradientBoostingClassifier(random_state=42, **params),
45
+ },
46
+ "regression": {
47
+ "random_forest": lambda **params: RandomForestRegressor(random_state=42, **params),
48
+ "linear_regression": lambda **params: LinearRegression(**params),
49
+ "gradient_boosting": lambda **params: GradientBoostingRegressor(random_state=42, **params),
50
+ },
51
+ }
52
+
53
+ _DEFAULT_SCORING = {"classification": "accuracy", "regression": "r2"}
54
+
55
+
56
+ @dataclass(frozen=True)
57
+ class ModelResult:
58
+ model: object
59
+ metrics: dict
60
+ train_index: list
61
+ test_index: list
62
+ scoring: str | None = None
63
+ fold_scores: list[float] | None = None
64
+ cv_splits: list[tuple[list, list]] | None = None
65
+ oof_predictions: pd.Series | None = None
66
+ oof_probabilities: pd.DataFrame | None = None
67
+
68
+
69
+ # @id CODE-AIDS-119
70
+ # @implements REQ-AIDS-074
71
+ # @design DES-AIDS-062
72
+ def build_cv_splits(
73
+ df: pd.DataFrame,
74
+ target: str,
75
+ model_type: str,
76
+ cv_strategy: str | None,
77
+ n_splits: int,
78
+ random_state: int,
79
+ groups: str | Sequence[Any] | pd.Series | None = None,
80
+ cv_splits: Sequence[tuple[Sequence[Any], Sequence[Any]]] | None = None,
81
+ ) -> list[tuple[list, list]] | None:
82
+ """Build or normalize reusable cross-validation splits."""
83
+ _validate_cv_dataframe_index(df)
84
+ if cv_splits is not None:
85
+ normalized = [(list(train_idx), list(test_idx)) for train_idx, test_idx in cv_splits]
86
+ if not normalized:
87
+ raise ValueError("cv_splits must not be empty")
88
+ _validate_cv_splits(df.index, normalized)
89
+ return normalized
90
+ if cv_strategy is None:
91
+ return None
92
+ if n_splits < 2:
93
+ raise ValueError("n_splits must be at least 2")
94
+
95
+ features = df.drop(columns=[target])
96
+ labels = df[target]
97
+ if cv_strategy == "StratifiedKFold":
98
+ if model_type != "classification":
99
+ raise ValueError("StratifiedKFold is supported only for classification")
100
+ splitter = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=random_state)
101
+ split_iter = splitter.split(features, labels)
102
+ elif cv_strategy == "KFold":
103
+ splitter = KFold(n_splits=n_splits, shuffle=True, random_state=random_state)
104
+ split_iter = splitter.split(features, labels)
105
+ elif cv_strategy == "GroupKFold":
106
+ group_values = _normalize_groups(df, groups)
107
+ splitter = GroupKFold(n_splits=n_splits)
108
+ split_iter = splitter.split(features, labels, group_values)
109
+ else:
110
+ raise ValueError(f"Unsupported cv_strategy: {cv_strategy!r}")
111
+
112
+ normalized = [
113
+ (df.index[train_positions].tolist(), df.index[test_positions].tolist())
114
+ for train_positions, test_positions in split_iter
115
+ ]
116
+ _validate_cv_splits(df.index, normalized)
117
+ return normalized
118
+
119
+
120
+ # @id CODE-AIDS-120
121
+ # @implements REQ-AIDS-078
122
+ # @design DES-AIDS-066
123
+ def resolve_estimator(
124
+ model_type: str,
125
+ model_name: str | None = None,
126
+ estimator: object | Callable[..., object] | None = None,
127
+ model_params: dict[str, Any] | None = None,
128
+ ) -> object:
129
+ """Resolve a built-in model or external estimator into a fresh estimator."""
130
+ model_params = dict(model_params or {})
131
+ if estimator is not None:
132
+ return _materialize_external_estimator(estimator, model_params)
133
+
134
+ builders = MODEL_BUILDERS.get(model_type)
135
+ if builders is None:
136
+ raise ValueError(f"Unsupported model_type: {model_type!r}")
137
+ resolved_model_name = model_name or next(iter(builders))
138
+ if resolved_model_name not in builders:
139
+ raise ValueError(f"Unsupported model_name {resolved_model_name!r} for {model_type!r}")
140
+ return builders[resolved_model_name](**model_params)
141
+
142
+
143
+ # @id CODE-AIDS-121
144
+ # @implements REQ-AIDS-075
145
+ # @design DES-AIDS-063
146
+ def compute_score(
147
+ model_type: str,
148
+ scoring: str,
149
+ y_true: pd.Series,
150
+ predictions: pd.Series,
151
+ probabilities: pd.DataFrame | None = None,
152
+ ) -> float:
153
+ """Compute the requested score from held-out predictions."""
154
+ if model_type == "classification":
155
+ if scoring == "accuracy":
156
+ return float(accuracy_score(y_true, predictions))
157
+ if scoring == "precision":
158
+ return float(precision_score(y_true, predictions, average="macro", zero_division=0))
159
+ if scoring == "recall":
160
+ return float(recall_score(y_true, predictions, average="macro", zero_division=0))
161
+ if scoring == "roc_auc":
162
+ probabilities = _require_probabilities(probabilities, scoring)
163
+ if probabilities.shape[1] == 2:
164
+ return float(roc_auc_score(y_true, probabilities.iloc[:, -1]))
165
+ return float(roc_auc_score(y_true, probabilities, multi_class="ovr"))
166
+ if scoring == "log_loss":
167
+ probabilities = _require_probabilities(probabilities, scoring)
168
+ return float(log_loss(y_true, probabilities, labels=list(probabilities.columns)))
169
+ raise ValueError(f"Unsupported scoring {scoring!r} for classification")
170
+
171
+ if scoring == "r2":
172
+ return float(r2_score(y_true, predictions))
173
+ if scoring == "rmse":
174
+ return float(np.sqrt(mean_squared_error(y_true, predictions)))
175
+ raise ValueError(f"Unsupported scoring {scoring!r} for regression")
176
+
177
+
178
+ def _metric_direction(model_type: str, scoring: str | None) -> bool:
179
+ selected = scoring or _DEFAULT_SCORING[model_type]
180
+ return selected not in {"log_loss", "rmse"}
181
+
182
+
183
+ def _materialize_external_estimator(
184
+ estimator: object | Callable[..., object], model_params: dict[str, Any]
185
+ ) -> object:
186
+ if hasattr(estimator, "fit") and hasattr(estimator, "predict"):
187
+ resolved = clone(estimator)
188
+ if model_params:
189
+ if not hasattr(resolved, "set_params"):
190
+ raise ValueError("External estimator does not support parameter overrides")
191
+ resolved = resolved.set_params(**model_params)
192
+ return resolved
193
+ if callable(estimator):
194
+ resolved = estimator(**model_params)
195
+ if not hasattr(resolved, "fit") or not hasattr(resolved, "predict"):
196
+ raise ValueError("Resolved estimator must implement fit and predict")
197
+ return resolved
198
+ raise ValueError("External estimator must be cloneable or callable")
199
+
200
+
201
+ def _normalize_groups(
202
+ df: pd.DataFrame, groups: str | Sequence[Any] | pd.Series | None
203
+ ) -> pd.Series:
204
+ if groups is None:
205
+ raise ValueError("groups are required when cv_strategy='GroupKFold'")
206
+ if isinstance(groups, str):
207
+ if groups not in df.columns:
208
+ raise ValueError(f"Unknown groups column: {groups!r}")
209
+ return df[groups]
210
+ group_series = pd.Series(groups, index=df.index)
211
+ if len(group_series) != len(df.index):
212
+ raise ValueError("groups must have one value per input row")
213
+ return group_series
214
+
215
+
216
+ def _require_probabilities(probabilities: pd.DataFrame | None, scoring: str) -> pd.DataFrame:
217
+ if probabilities is None:
218
+ raise ValueError(f"scoring={scoring!r} requires an estimator with predict_proba")
219
+ return probabilities
220
+
221
+
222
+ def _predict_probabilities(model: object, x_test: pd.DataFrame) -> pd.DataFrame | None:
223
+ if not hasattr(model, "predict_proba"):
224
+ return None
225
+ probability_values = model.predict_proba(x_test)
226
+ classes = list(getattr(model, "classes_", range(probability_values.shape[1])))
227
+ return pd.DataFrame(probability_values, index=x_test.index, columns=classes)
228
+
229
+
230
+ def _validate_cv_dataframe_index(df: pd.DataFrame) -> None:
231
+ if not df.index.is_unique:
232
+ raise ValueError("Cross-validation requires a DataFrame with unique index labels")
233
+
234
+
235
+ def _validate_cv_splits(
236
+ index: pd.Index, cv_splits: Sequence[tuple[Sequence[Any], Sequence[Any]]]
237
+ ) -> None:
238
+ allowed_labels = set(index.tolist())
239
+ test_counts: Counter[Any] = Counter()
240
+
241
+ for fold_number, (train_idx, test_idx) in enumerate(cv_splits, start=1):
242
+ train_labels = list(train_idx)
243
+ test_labels = list(test_idx)
244
+ unknown_labels = (set(train_labels) | set(test_labels)) - allowed_labels
245
+ if unknown_labels:
246
+ raise ValueError(
247
+ "cv_splits fold "
248
+ f"{fold_number} contains unknown index labels: {sorted(unknown_labels)!r}"
249
+ )
250
+ overlap = set(train_labels) & set(test_labels)
251
+ if overlap:
252
+ raise ValueError(
253
+ f"cv_splits fold {fold_number} has train/test overlap: {sorted(overlap)!r}"
254
+ )
255
+ duplicate_test_labels = [
256
+ label for label, count in Counter(test_labels).items() if count > 1
257
+ ]
258
+ if duplicate_test_labels:
259
+ raise ValueError(
260
+ "cv_splits fold "
261
+ f"{fold_number} repeats test index labels: {sorted(duplicate_test_labels)!r}"
262
+ )
263
+ test_counts.update(test_labels)
264
+
265
+ missing_test_labels = [label for label in index if test_counts[label] == 0]
266
+ duplicate_test_coverage = [label for label, count in test_counts.items() if count > 1]
267
+ if missing_test_labels:
268
+ raise ValueError(
269
+ "cv_splits must assign every row to exactly one test fold; "
270
+ f"missing test coverage for: {missing_test_labels!r}"
271
+ )
272
+ if duplicate_test_coverage:
273
+ raise ValueError(
274
+ "cv_splits must assign every row to exactly one test fold; "
275
+ f"duplicate test coverage for: {sorted(duplicate_test_coverage)!r}"
276
+ )
277
+
278
+
279
+ def _classification_metrics(
280
+ y_true: pd.Series,
281
+ predictions: pd.Series,
282
+ scoring: str | None,
283
+ probabilities: pd.DataFrame | None,
284
+ ) -> dict[str, float]:
285
+ metrics = {
286
+ "accuracy": float(accuracy_score(y_true, predictions)),
287
+ "precision": float(precision_score(y_true, predictions, average="macro", zero_division=0)),
288
+ "recall": float(recall_score(y_true, predictions, average="macro", zero_division=0)),
289
+ }
290
+ if scoring and scoring not in metrics:
291
+ metrics[scoring] = compute_score(
292
+ "classification", scoring, y_true, predictions, probabilities
293
+ )
294
+ return metrics
295
+
296
+
297
+ def _regression_metrics(
298
+ y_true: pd.Series, predictions: pd.Series, scoring: str | None
299
+ ) -> dict[str, float]:
300
+ metrics = {
301
+ "rmse": float(np.sqrt(mean_squared_error(y_true, predictions))),
302
+ "r2": float(r2_score(y_true, predictions)),
303
+ }
304
+ if scoring and scoring not in metrics:
305
+ metrics[scoring] = compute_score("regression", scoring, y_true, predictions)
306
+ return metrics
307
+
308
+
309
+ # @id CODE-AIDS-008
310
+ # @implements REQ-AIDS-008
311
+ # @design DES-AIDS-012
312
+ def train_model(
313
+ df: pd.DataFrame,
314
+ target: str,
315
+ model_type: str = "classification",
316
+ model_name: str = "random_forest",
317
+ test_size: float = 0.2,
318
+ random_state: int = 42,
319
+ scoring: str | None = None,
320
+ cv_strategy: str | None = None,
321
+ n_splits: int = 5,
322
+ groups: str | Sequence[Any] | pd.Series | None = None,
323
+ cv_splits: Sequence[tuple[Sequence[Any], Sequence[Any]]] | None = None,
324
+ estimator: object | Callable[..., object] | None = None,
325
+ **model_params,
326
+ ) -> ModelResult:
327
+ """Train a classification or regression model on ``df``.
328
+
329
+ Defaults to the legacy single holdout split. When cross-validation is
330
+ requested, returns fold scores and out-of-fold artifacts aligned to the
331
+ original row order.
332
+ """
333
+ if model_type not in MODEL_BUILDERS:
334
+ raise ValueError(f"Unsupported model_type: {model_type!r}")
335
+
336
+ feature_columns = [c for c in df.columns if c != target]
337
+ x = df[feature_columns]
338
+ y = df[target]
339
+ selected_scoring = scoring or _DEFAULT_SCORING[model_type]
340
+
341
+ # @id CODE-AIDS-122
342
+ # @implements REQ-AIDS-074 REQ-AIDS-075 REQ-AIDS-078
343
+ # @design DES-AIDS-063
344
+ normalized_cv_splits = build_cv_splits(
345
+ df=df,
346
+ target=target,
347
+ model_type=model_type,
348
+ cv_strategy=cv_strategy,
349
+ n_splits=n_splits,
350
+ random_state=random_state,
351
+ groups=groups,
352
+ cv_splits=cv_splits,
353
+ )
354
+ if normalized_cv_splits is None:
355
+ train_idx, test_idx = train_test_split(
356
+ df.index, test_size=test_size, random_state=random_state
357
+ )
358
+ model = resolve_estimator(
359
+ model_type=model_type,
360
+ model_name=model_name,
361
+ estimator=estimator,
362
+ model_params=model_params,
363
+ )
364
+ model.fit(x.loc[train_idx], y.loc[train_idx])
365
+ predictions = pd.Series(model.predict(x.loc[test_idx]), index=test_idx)
366
+ probabilities = (
367
+ _predict_probabilities(model, x.loc[test_idx])
368
+ if model_type == "classification"
369
+ else None
370
+ )
371
+ metrics = (
372
+ _classification_metrics(y.loc[test_idx], predictions, scoring, probabilities)
373
+ if model_type == "classification"
374
+ else _regression_metrics(y.loc[test_idx], predictions, scoring)
375
+ )
376
+ return ModelResult(
377
+ model=model,
378
+ metrics=metrics,
379
+ train_index=list(train_idx),
380
+ test_index=list(test_idx),
381
+ scoring=scoring,
382
+ )
383
+
384
+ oof_predictions = pd.Series(index=df.index, dtype=object)
385
+ oof_probabilities: pd.DataFrame | None = None
386
+ fold_scores: list[float] = []
387
+
388
+ for train_idx, test_idx in normalized_cv_splits:
389
+ fold_model = resolve_estimator(
390
+ model_type=model_type,
391
+ model_name=model_name,
392
+ estimator=estimator,
393
+ model_params=model_params,
394
+ )
395
+ x_train = x.loc[train_idx]
396
+ x_test = x.loc[test_idx]
397
+ y_train = y.loc[train_idx]
398
+ y_test = y.loc[test_idx]
399
+ fold_model.fit(x_train, y_train)
400
+
401
+ fold_predictions = pd.Series(fold_model.predict(x_test), index=test_idx)
402
+ oof_predictions.loc[test_idx] = fold_predictions
403
+ fold_probabilities = (
404
+ _predict_probabilities(fold_model, x_test) if model_type == "classification" else None
405
+ )
406
+ if fold_probabilities is not None:
407
+ if oof_probabilities is None:
408
+ oof_probabilities = pd.DataFrame(
409
+ index=df.index, columns=fold_probabilities.columns, dtype=float
410
+ )
411
+ oof_probabilities.loc[test_idx, fold_probabilities.columns] = fold_probabilities
412
+
413
+ fold_scores.append(
414
+ compute_score(
415
+ model_type, selected_scoring, y_test, fold_predictions, fold_probabilities
416
+ )
417
+ )
418
+
419
+ final_model = resolve_estimator(
420
+ model_type=model_type,
421
+ model_name=model_name,
422
+ estimator=estimator,
423
+ model_params=model_params,
424
+ )
425
+ final_model.fit(x, y)
426
+
427
+ cast_predictions = oof_predictions.astype(y.dtype)
428
+ metrics = (
429
+ _classification_metrics(y, cast_predictions, scoring, oof_probabilities)
430
+ if model_type == "classification"
431
+ else _regression_metrics(y, cast_predictions.astype(float), scoring)
432
+ )
433
+ metrics[selected_scoring] = compute_score(
434
+ model_type,
435
+ selected_scoring,
436
+ y,
437
+ cast_predictions if model_type == "classification" else cast_predictions.astype(float),
438
+ oof_probabilities,
439
+ )
440
+
441
+ return ModelResult(
442
+ model=final_model,
443
+ metrics=metrics,
444
+ train_index=list(normalized_cv_splits[0][0]),
445
+ test_index=list(normalized_cv_splits[0][1]),
446
+ scoring=selected_scoring,
447
+ fold_scores=fold_scores,
448
+ cv_splits=normalized_cv_splits,
449
+ oof_predictions=cast_predictions,
450
+ oof_probabilities=oof_probabilities,
451
+ )
@@ -0,0 +1,104 @@
1
+ """Hyperparameter tuning and model comparison.
2
+
3
+ Implements DES-AIDS-017 (REQ-AIDS-019): evaluates multiple parameter sets
4
+ or model candidates against the supervised modeling interface and reports
5
+ the best-performing configuration with its metric.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from collections.abc import Sequence
11
+ from dataclasses import dataclass
12
+ from typing import Any
13
+
14
+ import pandas as pd
15
+
16
+ from ai_data_scientist.ml_modeling import (
17
+ _DEFAULT_SCORING,
18
+ _metric_direction,
19
+ build_cv_splits,
20
+ train_model,
21
+ )
22
+
23
+
24
+ @dataclass(frozen=True)
25
+ class TuningResult:
26
+ best_params: dict
27
+ best_metric: float
28
+ all_candidates: list
29
+ scoring: str | None = None
30
+ cv_splits: list[tuple[list, list]] | None = None
31
+ best_result: object | None = None
32
+
33
+
34
+ # @id CODE-AIDS-019
35
+ # @implements REQ-AIDS-019
36
+ # @design DES-AIDS-017
37
+ def tune_or_compare(
38
+ df: pd.DataFrame,
39
+ target: str,
40
+ grid: list,
41
+ model_type: str = "classification",
42
+ scoring: str | None = None,
43
+ cv_strategy: str | None = None,
44
+ n_splits: int = 5,
45
+ random_state: int = 42,
46
+ groups: str | Sequence[Any] | pd.Series | None = None,
47
+ cv_splits: Sequence[tuple[Sequence[Any], Sequence[Any]]] | None = None,
48
+ ) -> TuningResult:
49
+ """Evaluate each parameter set in ``grid`` and report the best one."""
50
+ metric_name = scoring or _DEFAULT_SCORING[model_type]
51
+
52
+ # @id CODE-AIDS-123
53
+ # @implements REQ-AIDS-076 REQ-AIDS-078
54
+ # @design DES-AIDS-064
55
+ shared_cv_splits = build_cv_splits(
56
+ df=df,
57
+ target=target,
58
+ model_type=model_type,
59
+ cv_strategy=cv_strategy,
60
+ n_splits=n_splits,
61
+ random_state=random_state,
62
+ groups=groups,
63
+ cv_splits=cv_splits,
64
+ )
65
+ candidates = []
66
+ for params in grid:
67
+ params_copy = dict(params)
68
+ result = train_model(
69
+ df,
70
+ target=target,
71
+ model_type=model_type,
72
+ scoring=scoring,
73
+ cv_strategy=cv_strategy,
74
+ n_splits=n_splits,
75
+ random_state=random_state,
76
+ groups=groups,
77
+ cv_splits=shared_cv_splits,
78
+ **params_copy,
79
+ )
80
+ candidates.append(
81
+ {
82
+ "params": params_copy,
83
+ "metric": result.metrics[metric_name],
84
+ "fold_scores": result.fold_scores,
85
+ "result": result,
86
+ "cv_splits": result.cv_splits,
87
+ }
88
+ )
89
+
90
+ choose_max = _metric_direction(model_type, scoring)
91
+ best = (
92
+ max(candidates, key=lambda candidate: candidate["metric"])
93
+ if choose_max
94
+ else min(candidates, key=lambda candidate: candidate["metric"])
95
+ )
96
+
97
+ return TuningResult(
98
+ best_params=best["params"],
99
+ best_metric=best["metric"],
100
+ all_candidates=candidates,
101
+ scoring=metric_name,
102
+ cv_splits=shared_cv_splits,
103
+ best_result=best["result"],
104
+ )