Spaces:
Running
Running
| import { useEffect, useState, useRef } from "react"; | |
| import * as Comlink from "comlink"; | |
| import type { ModelSettings, DataPoints, DecisionBoundarySettings, DecisionBoundaryResult, CsvSettings } from "./types.ts"; | |
| import { OKABE_ITO_COLORS, LabelColors } from "./colors.ts"; | |
| import usePyodideBackend from "./usePyodideBackend.ts"; | |
| const DEFAULT_CSV_SETTINGS: CsvSettings = { | |
| inputColumns: "1,2", | |
| outputColumn: "3", | |
| normalizerType: "None", | |
| normalNoiseStd: "0.0", | |
| projectionType: "Coordinates", | |
| x1Column: "1", | |
| x2Column: "2", | |
| }; | |
| const DEFAULT_RESOLUTION = 500; | |
| const DEFAULT_X_RANGE: [number, number] = [-10, 10]; | |
| const DEFAULT_Y_RANGE: [number, number] = [-10, 10]; | |
| export default function useAppLogic() { | |
| const { getBackend, backendReady } = usePyodideBackend(); | |
| const [dataPoints, setDataPoints] = useState<DataPoints>({ | |
| xPoints: [], | |
| yPoints: [], | |
| labels: [], | |
| }); | |
| const labelColors = new LabelColors(dataPoints.labels); | |
| const [canAddPoints, setCanAddPoints] = useState<boolean>(true); | |
| const [pointLabel, setPointLabel] = useState<string>(OKABE_ITO_COLORS[0].name); | |
| const [resolution, _] = useState<number>(DEFAULT_RESOLUTION); // todo user defined | |
| const xRangeRef = useRef<[number, number]>(DEFAULT_X_RANGE); | |
| const yRangeRef = useRef<[number, number]>(DEFAULT_Y_RANGE); | |
| const [rangeVersion, setRangeVersion] = useState<number>(0); | |
| function handleRangeChange(xRange: [number, number], yRange: [number, number]) { | |
| xRangeRef.current = xRange; | |
| yRangeRef.current = yRange; | |
| updateDecisionBoundary(); | |
| } | |
| function handleRangeChangeWithRender(xRange: [number, number], yRange: [number, number]) { | |
| xRangeRef.current = xRange; | |
| yRangeRef.current = yRange; | |
| setRangeVersion((prev) => prev + 1); | |
| console.log(xRangeRef.current, yRangeRef.current); | |
| } | |
| const [modelSettings, setModelSettings] = useState<ModelSettings>({ | |
| type: "LogisticRegression", | |
| arguments: "", | |
| }); | |
| async function handleModelSettingsChange(settings: ModelSettings) { | |
| setModelSettings(settings); | |
| if (!backendReady) { | |
| return; | |
| } | |
| await getBackend().setModelConfig(settings); | |
| await updateDecisionBoundary(); | |
| } | |
| const [csvSettings, setCsvSettings] = useState<CsvSettings>(DEFAULT_CSV_SETTINGS); | |
| const [csvError, setCsvError] = useState<string | null>(null); | |
| const [decisionBoundaryResult, setDecisionBoundaryResult] = useState<DecisionBoundaryResult | null>(null); | |
| // this is needed for svg export | |
| const plotHandle = useRef<Plotly.PlotlyHTMLElement | null>(null); | |
| async function handleAddDataPoint(x: number, y: number, label: string) { | |
| const next = { | |
| xPoints: [...dataPoints.xPoints, x], | |
| yPoints: [...dataPoints.yPoints, y], | |
| labels: [...dataPoints.labels, label], | |
| } | |
| setDataPoints(next); | |
| if (!backendReady) { | |
| return; | |
| } | |
| await getBackend().setDataset2d(next); | |
| await updateDecisionBoundary(); | |
| } | |
| async function handleUndoDataPoint() { | |
| const next = { | |
| xPoints: dataPoints.xPoints.slice(0, -1), | |
| yPoints: dataPoints.yPoints.slice(0, -1), | |
| labels: dataPoints.labels.slice(0, -1), | |
| } | |
| setDataPoints(next); | |
| if (!backendReady) { | |
| return; | |
| } | |
| await getBackend().setDataset2d(next); | |
| await updateDecisionBoundary(); | |
| } | |
| async function handleClearDataPoints() { | |
| const next = { | |
| xPoints: [], | |
| yPoints: [], | |
| labels: [], | |
| } | |
| setDataPoints(next); | |
| if (!backendReady) { | |
| return; | |
| } | |
| await getBackend().setDataset2d(next); | |
| await updateDecisionBoundary(); | |
| } | |
| async function updateDecisionBoundary() { | |
| if (!backendReady) { | |
| return; | |
| } | |
| const settings: DecisionBoundarySettings = { | |
| xmin: xRangeRef.current[0], | |
| xmax: xRangeRef.current[1], | |
| ymin: yRangeRef.current[0], | |
| ymax: yRangeRef.current[1], | |
| resolution: resolution, | |
| }; | |
| const result = await getBackend().getDecisionBoundary(settings); | |
| setDecisionBoundaryResult(result); | |
| } | |
| async function handleGetDecisionBoundary() { | |
| if (!backendReady) { | |
| return; | |
| } | |
| await getBackend().buildModel(); | |
| await updateDecisionBoundary(); | |
| } | |
| async function handleCsvUpload(file: File, settings?: CsvSettings) { | |
| const activeCsvSettings = settings ?? csvSettings; | |
| if (settings) { | |
| setCsvSettings(settings); | |
| } | |
| if (!backendReady) { | |
| alert("Python backend is not ready yet."); | |
| return; | |
| } | |
| const buffer = await file.arrayBuffer(); | |
| console.log(buffer); | |
| const result = await getBackend().setDatasetCsv( | |
| Comlink.transfer(buffer, [buffer]), | |
| activeCsvSettings, | |
| ); | |
| if (result.status === "OK") { | |
| setDataPoints(result.dataPoints); | |
| if (result.xRange && result.yRange) { | |
| handleRangeChangeWithRender(result.xRange, result.yRange); | |
| } | |
| setCsvError(null); | |
| } else if (result.status === "CSV_ERROR") { | |
| setCsvError(result.message); | |
| } | |
| await updateDecisionBoundary(); | |
| } | |
| async function handleCsvSettingsChange(settings: CsvSettings) { | |
| if (!backendReady) { | |
| return; | |
| } | |
| setCsvSettings(settings); | |
| const result = await getBackend().setCsvSettings(settings); | |
| if (result.status === "OK") { | |
| setDataPoints(result.dataPoints); | |
| if (result.xRange && result.yRange) { | |
| handleRangeChangeWithRender(result.xRange, result.yRange); | |
| } | |
| setCsvError(null); | |
| } else if (result.status === "CSV_ERROR") { | |
| setCsvError(result.message); | |
| } | |
| await updateDecisionBoundary(); | |
| } | |
| function downloadCsv(filename: string, csvText: string) { | |
| const blob = new Blob([csvText], { type: "text/csv;charset=utf-8" }); | |
| const url = URL.createObjectURL(blob); | |
| const a = document.createElement("a"); | |
| a.href = url; | |
| a.download = filename; | |
| document.body.appendChild(a); | |
| a.click(); | |
| a.remove(); | |
| URL.revokeObjectURL(url); | |
| } | |
| async function getCsvDataset() { | |
| if (!backendReady) { | |
| alert("Python backend is not ready yet."); | |
| return; | |
| } | |
| const csvData = await getBackend().getDatasetCsv(); | |
| downloadCsv("dataset.csv", csvData); | |
| } | |
| function getPlotSvg() { | |
| if (!plotHandle.current) { | |
| return; | |
| } | |
| (plotHandle.current as any)?.downloadSvg() | |
| } | |
| console.log("rendering app"); | |
| useEffect(() => { | |
| if (!backendReady) { | |
| return; | |
| } | |
| (async () => { | |
| await getBackend().setDataset2d(dataPoints); | |
| await getBackend().setModelConfig(modelSettings); | |
| await updateDecisionBoundary(); | |
| })(); | |
| }, [backendReady, getBackend]); | |
| return { | |
| backendReady, | |
| dataPoints, | |
| labelColors, | |
| canAddPoints, | |
| setCanAddPoints, | |
| pointLabel, | |
| setPointLabel, | |
| xRangeRef, | |
| yRangeRef, | |
| handleRangeChange, | |
| rangeVersion, | |
| modelSettings, | |
| handleModelSettingsChange, | |
| csvSettings, | |
| handleCsvSettingsChange, | |
| csvError, | |
| decisionBoundaryResult, | |
| handleAddDataPoint, | |
| handleUndoDataPoint, | |
| handleClearDataPoints, | |
| handleGetDecisionBoundary, | |
| handleCsvUpload, | |
| getCsvDataset, | |
| plotHandle, | |
| getPlotSvg, | |
| } | |
| } | |