decision_boundary / backend /src /manager.py
Joel Woodfield
Fix bug where projector was not reset when dataset changes
290a07c
Raw
History Blame Contribute Delete
8.78 kB
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