@@ -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+
47127def 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 ,
0 commit comments