Skip to content

Commit f0de28c

Browse files
fix: scitype-correct export_code examples, is_pipeline, dataset, loaded models (#534)
Fixes #534. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
1 parent 4d809ea commit f0de28c

2 files changed

Lines changed: 226 additions & 34 deletions

File tree

src/sktime_mcp/tools/codegen.py

Lines changed: 111 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,86 @@ def _is_valid_var_name(var_name: str) -> bool:
4444
return isinstance(var_name, str) and var_name.isidentifier() and not keyword.iskeyword(var_name)
4545

4646

47+
def _loader_for(dataset: str, demo_datasets: dict) -> tuple[str, str]:
48+
"""Return (module, func) for a demo dataset name, defaulting to load_airline."""
49+
if dataset in demo_datasets:
50+
module_path = demo_datasets[dataset]
51+
module, func = module_path.rsplit(".", 1)
52+
return module, func
53+
return "sktime.datasets", "load_airline"
54+
55+
56+
def _fit_example(
57+
var_name: str,
58+
obj_type: str,
59+
dataset: str | None,
60+
handle_info: Any,
61+
demo_datasets: dict,
62+
) -> str:
63+
"""Build a runnable fit/predict example matching the estimator's scitype.
64+
65+
A forecaster-shaped example (`fit(y)` / `predict(fh)`) is wrong for
66+
transformers, splitters, and classifiers, which raise AttributeError when
67+
the generated code runs (BUG-03).
68+
"""
69+
if obj_type in ("classifier", "regressor"):
70+
# Panel X + label/target y — use a classification demo dataset.
71+
ds = dataset or "arrow_head"
72+
module, func = _loader_for(ds, demo_datasets)
73+
verb = "class" if obj_type == "classifier" else "value"
74+
return f"""
75+
76+
# Example usage ({obj_type}):
77+
from {module} import {func}
78+
X, y = {func}(return_X_y=True)
79+
80+
{var_name}.fit(X, y)
81+
predictions = {var_name}.predict(X) # predicted {verb} per instance
82+
print(predictions)
83+
"""
84+
85+
if obj_type == "transformer":
86+
ds = dataset or handle_info.metadata.get("training_dataset") or "airline"
87+
module, func = _loader_for(ds, demo_datasets)
88+
return f"""
89+
90+
# Example usage (transformer):
91+
from {module} import {func}
92+
y = {func}()
93+
94+
y_transformed = {var_name}.fit_transform(y)
95+
print(y_transformed)
96+
"""
97+
98+
if obj_type == "splitter":
99+
ds = dataset or "airline"
100+
module, func = _loader_for(ds, demo_datasets)
101+
return f"""
102+
103+
# Example usage (splitter):
104+
from {module} import {func}
105+
y = {func}()
106+
107+
for train_idx, test_idx in {var_name}.split(y):
108+
print("train:", train_idx, "test:", test_idx)
109+
"""
110+
111+
# Default: forecaster.
112+
ds = dataset or handle_info.metadata.get("training_dataset") or "airline"
113+
module, func = _loader_for(ds, demo_datasets)
114+
return f"""
115+
116+
# Example usage (forecaster):
117+
from {module} import {func}
118+
y = {func}()
119+
120+
{var_name}.fit(y)
121+
fh = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12] # 12-step ahead forecast
122+
predictions = {var_name}.predict(fh=fh)
123+
print(predictions)
124+
"""
125+
126+
47127
def export_code_tool(
48128
handle: str,
49129
var_name: str = "model",
@@ -96,45 +176,42 @@ def export_code_tool(
96176

97177
estimator_name = handle_info.estimator_name
98178
params = handle_info.params
99-
100179
spec = params.get("spec")
101-
if not spec:
102-
return {"success": False, "error": "No craft spec found in handle parameters."}
103180

104-
is_pipeline = "*" in spec or "Pipeline" in spec or "[" in spec
105-
code = f"from sktime.registry import craft\n\n{var_name} = craft({_format_value(spec)})"
181+
instance = handle_manager.get_instance(handle)
182+
get_tag = getattr(instance, "get_class_tag", None)
183+
obj_type = get_tag("object_type", "") if callable(get_tag) else ""
184+
185+
# is_pipeline from the instance, not a spec substring — "[" in spec
186+
# false-positived on any list argument (BUG-04).
187+
is_pipeline = bool(spec and "*" in spec) or hasattr(instance, "steps")
188+
189+
if spec:
190+
code = f"from sktime.registry import craft\n\n{var_name} = craft({_format_value(spec)})"
191+
elif handle_info.metadata.get("source") == "loaded" and handle_info.metadata.get("path"):
192+
# Loaded models carry no craft spec; emit a load_model snippet instead of
193+
# failing with "No craft spec found" (NB-17).
194+
model_path = handle_info.metadata["path"]
195+
code = (
196+
"from sktime.utils.mlflow_sktime import load_model\n\n"
197+
f"{var_name} = load_model({_format_value(model_path)})"
198+
)
199+
else:
200+
return {"success": False, "error": "No craft spec found in handle parameters."}
106201

107-
# Optionally add fit/predict example
202+
# Optionally add a scitype-appropriate fit example (BUG-03).
108203
if include_fit_example:
109-
# Priority: explicit argument > dataset used during fit > "airline" fallback
110-
effective_dataset = dataset or handle_info.metadata.get("training_dataset") or "airline"
111-
# Resolve the dataset loader from discovered demo datasets
112204
demo_datasets = _get_demo_datasets()
113-
if effective_dataset in demo_datasets:
114-
module_path = demo_datasets[effective_dataset]
115-
module_parts = module_path.rsplit(".", 1)
116-
loader_module = module_parts[0]
117-
loader_func = module_parts[1]
118-
else:
119-
loader_module = "sktime.datasets"
120-
loader_func = "load_airline"
121-
122-
example_code = f"""
123-
124-
# Example usage:
125-
# Load data
126-
from {loader_module} import {loader_func}
127-
y = {loader_func}()
128-
129-
# Fit the model
130-
{var_name}.fit(y)
131-
132-
# Make predictions
133-
fh = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12] # 12-step ahead forecast
134-
predictions = {var_name}.predict(fh=fh)
135-
print(predictions)
136-
"""
137-
code += example_code
205+
if dataset is not None and dataset not in demo_datasets:
206+
return {
207+
"success": False,
208+
"error": (
209+
f"Unknown dataset '{dataset}' for the fit example. Use a demo dataset "
210+
"name (see list_available_data) or omit dataset to use a default."
211+
),
212+
}
213+
example = _fit_example(var_name, obj_type, dataset, handle_info, demo_datasets)
214+
code += example
138215

139216
return {
140217
"success": True,

tests/test_export_code_scitype.py

Lines changed: 115 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,115 @@
1+
"""export_code must generate correct, runnable code per scitype (#534).
2+
3+
- BUG-03: fit example was forecaster-shaped for every scitype (AttributeError
4+
when run for transformers/splitters/classifiers).
5+
- BUG-04: is_pipeline false-positived on any "[" in the spec.
6+
- NB-15: the dataset argument was ignored; unknown datasets silently fell back
7+
to airline.
8+
- NB-17: loaded models (empty params) failed with "No craft spec found".
9+
"""
10+
11+
import contextlib
12+
13+
import pytest
14+
15+
from sktime_mcp.runtime.handles import get_handle_manager
16+
from sktime_mcp.tools.codegen import export_code_tool
17+
from sktime_mcp.tools.instantiate import instantiate_tool
18+
19+
20+
def _handle(spec):
21+
res = instantiate_tool(spec=spec)
22+
assert res["success"], res
23+
return res["handle"]
24+
25+
26+
def _release(handle):
27+
with contextlib.suppress(KeyError):
28+
get_handle_manager().release_handle(handle)
29+
30+
31+
class TestIsPipeline:
32+
def test_list_arg_is_not_a_pipeline(self):
33+
h = _handle("SlidingWindowSplitter(window_length=24, fh=[1, 2, 3], step_length=12)")
34+
try:
35+
res = export_code_tool(h)
36+
assert res["success"]
37+
assert res["is_pipeline"] is False
38+
finally:
39+
_release(h)
40+
41+
def test_star_spec_is_a_pipeline(self):
42+
h = _handle("Deseasonalizer() * NaiveForecaster()")
43+
try:
44+
res = export_code_tool(h)
45+
assert res["success"]
46+
assert res["is_pipeline"] is True
47+
finally:
48+
_release(h)
49+
50+
51+
class TestScitypeExampleRuns:
52+
def _export_and_exec(self, spec):
53+
h = _handle(spec)
54+
try:
55+
res = export_code_tool(h, include_fit_example=True)
56+
assert res["success"], res
57+
compile(res["code"], "<export>", "exec")
58+
exec(res["code"], {}) # must run without AttributeError
59+
finally:
60+
_release(h)
61+
62+
def test_forecaster_example_runs(self):
63+
self._export_and_exec("NaiveForecaster(sp=12)")
64+
65+
def test_transformer_example_runs(self):
66+
self._export_and_exec("Deseasonalizer()")
67+
68+
def test_splitter_example_runs(self):
69+
self._export_and_exec("SlidingWindowSplitter(window_length=24, step_length=12)")
70+
71+
def test_classifier_example_runs(self):
72+
self._export_and_exec("KNeighborsTimeSeriesClassifier()")
73+
74+
75+
class TestLoadedModelExport:
76+
def test_loaded_model_emits_load_model_snippet(self):
77+
from sktime.forecasting.naive import NaiveForecaster
78+
79+
hm = get_handle_manager()
80+
# Simulate a load_model handle: no craft spec, metadata carries the path.
81+
handle = hm.create_handle(
82+
estimator_name="NaiveForecaster",
83+
instance=NaiveForecaster(),
84+
params={},
85+
metadata={"source": "loaded", "path": "/tmp/some_model_dir"},
86+
)
87+
try:
88+
res = export_code_tool(handle)
89+
assert res["success"], res
90+
assert "load_model" in res["code"]
91+
assert "/tmp/some_model_dir" in res["code"]
92+
compile(res["code"], "<export>", "exec")
93+
finally:
94+
_release(handle)
95+
96+
97+
class TestDatasetValidation:
98+
def test_unknown_dataset_rejected(self):
99+
h = _handle("NaiveForecaster()")
100+
try:
101+
res = export_code_tool(h, include_fit_example=True, dataset="not_a_dataset_zzz")
102+
assert not res["success"]
103+
assert "not_a_dataset_zzz" in res["error"]
104+
finally:
105+
_release(h)
106+
107+
def test_known_dataset_used(self):
108+
h = _handle("NaiveForecaster()")
109+
try:
110+
res = export_code_tool(h, include_fit_example=True, dataset="lynx")
111+
assert res["success"]
112+
assert "load_lynx" in res["code"]
113+
assert "load_airline" not in res["code"]
114+
finally:
115+
_release(h)

0 commit comments

Comments
 (0)