decision_boundary / frontends /react /src /useAppLogic.ts
Joel Woodfield
Auto rescale plot ranges when using csv or preset dataset
803a9d0
Raw
History Blame Contribute Delete
7.13 kB
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,
}
}