Spaces:
Running
Running
| import csv | |
| from dataclasses import asdict | |
| from sklearn.exceptions import NotFittedError | |
| from sklearn.utils.validation import check_is_fitted | |
| from logic import * | |
| class Manager: | |
| def __init__(self): | |
| self._dataset: Dataset | None = None | |
| self._model: BaseEstimator | None = None | |
| self._model_config: ModelSettings | None = None | |
| self._csv: np.ndarray | None = None | |
| self._projector: Any | None = None | |
| self._preprocessor: Any | None = None | |
| def handle_set_dataset(self, dataset: dict) -> None: | |
| self._reset_model() | |
| self._projector = None | |
| self._preprocessor = None | |
| try: | |
| parsed_dataset = Dataset( | |
| inputs=np.array(dataset["inputs"]), | |
| outputs=np.array(dataset["outputs"]), | |
| ) | |
| except Exception as e: | |
| raise ValueError(f"Invalid dataset format: {e}") | |
| self._dataset = parsed_dataset | |
| def handle_set_dataset_csv(self, args: dict) -> dict: | |
| self._reset_model() | |
| self._projector = None | |
| self._preprocessor = None | |
| try: | |
| buffer = args["buffer"] | |
| settings = args["settings"] | |
| except KeyError as e: | |
| raise ValueError(f"Missing argument in handle_set_dataset_csv: {e}") | |
| try: | |
| csv = load_csv_js(buffer) | |
| except Exception as e: | |
| # user error | |
| message = f"Error loading dataset from CSV: {e}" | |
| return { "status": "USER_ERROR", "message": message } | |
| self._csv = csv | |
| result = self.handle_set_csv_settings(settings) | |
| return result | |
| def handle_set_csv_settings(self, settings: dict) -> dict: | |
| self._reset_model() | |
| # validate csv settings | |
| # USER_ERROR - program error (displays as alert) | |
| # CSV_ERROR - user input error (displays as message) | |
| if self._csv is None: | |
| return { "status": "CSV_ERROR", "message": "CSV data not loaded" } | |
| try: | |
| _ = settings["normalizerType"] | |
| _ = settings["normalNoiseStd"] | |
| _ = settings["inputColumns"] | |
| _ = settings["outputColumn"] | |
| _ = settings["projectionType"] | |
| except KeyError as e: | |
| raise ValueError(f"Missing CSV setting: {e}") | |
| if settings["normalizerType"] not in SUPPORTED_NORMALIZER_TYPES: | |
| raise ValueError(f"Unsupported normalizer type: {settings['normalizerType']}") | |
| if settings["projectionType"] not in SUPPORTED_PROJECTION_TYPES: | |
| raise ValueError(f"Unsupported projection type: {settings['projectionType']}") | |
| try: | |
| normal_noise_std = float(settings["normalNoiseStd"]) | |
| except ValueError: | |
| return { "status": "CSV_ERROR", "message": "normal_noise_std must be a number" } | |
| try: | |
| input_columns = parse_comma_separated_ints(settings["inputColumns"]) | |
| except ValueError: | |
| return { "status": "CSV_ERROR", "message": "inputColumns must be a comma-separated list of integers" } | |
| try: | |
| output_column = int(settings["outputColumn"]) | |
| except ValueError: | |
| return { "status": "CSV_ERROR", "message": "outputColumn must be an integer" } | |
| try: | |
| x1_column = int(settings.get("x1Column", "1")) | |
| except (ValueError, TypeError): | |
| return { "status": "CSV_ERROR", "message": "x1Column must be an integer" } | |
| try: | |
| x2_column = int(settings.get("x2Column", "2")) | |
| except (ValueError, TypeError): | |
| return { "status": "CSV_ERROR", "message": "x2Column must be an integer" } | |
| try: | |
| parsed_settings = CsvSettings( | |
| normalizer_type=settings["normalizerType"], | |
| normal_noise_std=normal_noise_std, | |
| input_columns=input_columns, | |
| output_column=output_column, | |
| projection_type=settings["projectionType"], | |
| x1_column=x1_column, | |
| x2_column=x2_column, | |
| ) | |
| except Exception as e: | |
| return { "status": "CSV_ERROR", "message": f"Invalid CSV settings format: {e}" } | |
| if len(parsed_settings.input_columns) < 2: | |
| return { "status": "CSV_ERROR", "message": "At least two input columns are required" } | |
| if max(parsed_settings.input_columns) > self._csv.shape[1] or min(parsed_settings.input_columns) < 1: | |
| return { "status": "CSV_ERROR", "message": "Input column index out of range" } | |
| if parsed_settings.output_column > self._csv.shape[1] or parsed_settings.output_column < 1: | |
| return { "status": "CSV_ERROR", "message": "Output column index out of range" } | |
| inputs = self._csv[:, [i - 1 for i in parsed_settings.input_columns]] | |
| outputs = self._csv[:, parsed_settings.output_column - 1] | |
| self._dataset = Dataset(inputs, outputs) | |
| self._preprocessor = init_preprocessor(self._dataset, parsed_settings) | |
| if self._preprocessor is not None: | |
| self._dataset.inputs = self._preprocessor.transform(self._dataset.inputs) | |
| self._projector = init_projector(self._dataset, parsed_settings) | |
| x1, x2, labels = project_dataset( | |
| self._dataset, | |
| self._projector, | |
| ) | |
| data_points = { "xPoints": x1, "yPoints": x2, "labels": labels } | |
| x1_range = np.array([np.min(x1), np.max(x1)]) | |
| x1_range += 0.1 * np.array([-1, 1]) * (x1_range[1] - x1_range[0]) | |
| x2_range = np.array([np.min(x2), np.max(x2)]) | |
| x2_range += 0.1 * np.array([-1, 1]) * (x2_range[1] - x2_range[0]) | |
| return { "status": "OK", "dataPoints": data_points, "xRange": x1_range.tolist(), "yRange": x2_range.tolist() } | |
| def handle_set_model_config(self, settings: dict) -> None: | |
| self._reset_model() | |
| try: | |
| parsed_settings = ModelSettings( | |
| type=settings["type"], | |
| arguments=settings["arguments"], | |
| ) | |
| except Exception as e: | |
| raise ValueError(f"Invalid model settings format: {e}") | |
| self._model_config = parsed_settings | |
| def handle_build_model(self) -> None: | |
| self._reset_model() | |
| if self._dataset is None: | |
| raise ValueError("Dataset not set") | |
| if self._model_config is None: | |
| raise ValueError("Model settings not set") | |
| try: | |
| self._model = init_model(self._model_config) | |
| except Exception as e: | |
| # user error | |
| message = f"Error initializing model: {e}" | |
| return { "status": "USER_ERROR", "message": message } | |
| if self._dataset.inputs.size > 0 and self._dataset.outputs.size > 0: | |
| try: | |
| train_model(self._model, self._dataset) | |
| except Exception as e: | |
| # user error | |
| self._reset_model() | |
| message = f"Error training model: {e}" | |
| return { "status": "USER_ERROR", "message": message } | |
| def handle_get_decision_boundary(self, settings: dict) -> dict: | |
| if self._model is None: | |
| return dict(x=[], y=[], labels=[]) | |
| try: | |
| check_is_fitted(self._model) | |
| except NotFittedError: | |
| return dict(x=[], y=[], labels=[]) | |
| try: | |
| parsed_settings = DecisionBoundarySettings( | |
| xmin=settings["xmin"], | |
| xmax=settings["xmax"], | |
| ymin=settings["ymin"], | |
| ymax=settings["ymax"], | |
| resolution=settings.get("resolution", 500), | |
| ) | |
| except Exception as e: | |
| raise ValueError(f"Invalid decision boundary settings format: {e}") | |
| try: | |
| result = get_decision_boundary_values( | |
| self._model, | |
| parsed_settings, | |
| self._projector.inverse_transform if self._projector else None, | |
| ) | |
| except Exception as e: | |
| message = f"Error computing decision boundary: {e}" | |
| return { "status": "USER_ERROR", "message": message } | |
| return asdict(result) | |
| def _reset_model(self) -> None: | |
| self._model = None | |
| print("Model reset") | |
| def handle_get_dataset_csv(self) -> str: | |
| if self._dataset is None: | |
| raise ValueError("Dataset not set") | |
| buf = io.StringIO() | |
| writer = csv.writer(buf) | |
| inputs = self._dataset.inputs | |
| outputs = self._dataset.outputs | |
| for row, label in zip(inputs, outputs): | |
| row_list = row.tolist() | |
| row_list.append(label) | |
| writer.writerow(row_list) | |
| csv_data = buf.getvalue() | |
| buf.close() | |
| return csv_data | |