From 58edf2bbf652c8aa4dd5e3de5e3f8d1b6d03be10 Mon Sep 17 00:00:00 2001 From: ahzero7d1 Date: Wed, 25 Mar 2026 13:29:37 +0100 Subject: [PATCH 1/6] Add invalid value check for tabular data --- .../src/components/dataset_input/validate.ts | 29 ++++++++++++++----- 1 file changed, 22 insertions(+), 7 deletions(-) diff --git a/webapp/src/components/dataset_input/validate.ts b/webapp/src/components/dataset_input/validate.ts index 38237bee3..886eb2711 100644 --- a/webapp/src/components/dataset_input/validate.ts +++ b/webapp/src/components/dataset_input/validate.ts @@ -2,15 +2,30 @@ import { Range, Set } from "immutable"; import type { LabeledDataset } from "./types"; +function isNaNValue(value: string | undefined): boolean{ + if (value === undefined) + return true; + + const trimmed = value.trim(); + return trimmed === "" || trimmed.toLowerCase() === "nan"; +} + export async function tabular( wantedColumns: Set, dataset: LabeledDataset["tabular"], ): Promise { - for await (const [columns, i] of dataset - .map((row) => Set(Object.keys(row))) - .zip(Range(1, Number.POSITIVE_INFINITY))) - if (!columns.isSuperset(wantedColumns)) - throw new Error( - `row ${i} is missing columns ${wantedColumns.subtract(columns).join(", ")}`, - ); + for await (const [row, i] of dataset + .zip(Range(1, Number.POSITIVE_INFINITY))){ + const columns = Set(Object.keys(row)); + + if (!columns.isSuperset(wantedColumns)) + throw new Error( + `row ${i} is missing columns ${wantedColumns.subtract(columns).join(", ")}`, + ); + + for (const col of wantedColumns){ + if (isNaNValue(row[col])) + throw new Error(`row ${i} column "${col}" contains NaN`); + } + } } From 8aaf95743ad3062c2a2d082a89636e3f3b24eedf Mon Sep 17 00:00:00 2001 From: ahzero7d1 Date: Mon, 30 Mar 2026 14:08:05 +0200 Subject: [PATCH 2/6] Add standardization for numerical features in tabular data --- discojs/src/index.ts | 1 + discojs/src/models/index.ts | 2 +- discojs/src/models/model.ts | 8 +++ discojs/src/models/tfjs.ts | 11 ++-- discojs/src/processing/index.ts | 17 +++++- discojs/src/processing/tabular.ts | 59 +++++++++++++++++++++ discojs/src/serialization/model.ts | 10 ++-- discojs/src/training/disco.ts | 85 +++++++++++++++++++++++++----- 8 files changed, 169 insertions(+), 24 deletions(-) diff --git a/discojs/src/index.ts b/discojs/src/index.ts index 04537e7c9..26fa556ec 100644 --- a/discojs/src/index.ts +++ b/discojs/src/index.ts @@ -17,6 +17,7 @@ export { EpochLogs, Tokenizer, ValidationMetrics, + ModelMetadata, } from "./models/index.js"; export * as models from './models/index.js' diff --git a/discojs/src/models/index.ts b/discojs/src/models/index.ts index 96d253f0b..d2497e894 100644 --- a/discojs/src/models/index.ts +++ b/discojs/src/models/index.ts @@ -1,4 +1,4 @@ -export { Model } from './model.js' +export { Model, ModelMetadata } from './model.js' export { BatchLogs, EpochLogs, ValidationMetrics } from "./logs.js"; export { Tokenizer } from "./tokenizer.js"; diff --git a/discojs/src/models/model.ts b/discojs/src/models/model.ts index dd7c0477c..9d617c621 100644 --- a/discojs/src/models/model.ts +++ b/discojs/src/models/model.ts @@ -7,6 +7,11 @@ import type { } from "../index.js"; import type { BatchLogs, EpochLogs } from "./logs.js"; +import type { StandardizationStats } from "../processing/tabular.js"; + +export type ModelMetadata = { + tabularStandardization?: StandardizationStats; +}; /** * Trainable predictor @@ -21,6 +26,9 @@ export abstract class Model implements Disposable { /** Set training state */ abstract set weights(ws: WeightsContainer); + /** Optional metadata for tabular task data standardization */ + metadata?: ModelMetadata; + /** * Improve predictor * diff --git a/discojs/src/models/tfjs.ts b/discojs/src/models/tfjs.ts index b60060f49..8a62dd5ca 100644 --- a/discojs/src/models/tfjs.ts +++ b/discojs/src/models/tfjs.ts @@ -12,17 +12,20 @@ import { import { BatchLogs } from './index.js' import { Model } from './index.js' import { EpochLogs } from './logs.js' +import { ModelMetadata } from "./model.js"; -type Serialized = [D, tf.io.ModelArtifacts]; +type Serialized = [D, tf.io.ModelArtifacts, ModelMetadata?]; /** TensorFlow JavaScript model with standard training */ export class TFJS extends Model { /** Wrap the given trainable model */ constructor ( public readonly datatype: D, - private readonly model: tf.LayersModel + private readonly model: tf.LayersModel, + metadata?: ModelMetadata, ) { super() + this.metadata = metadata; if (model.loss === undefined) { throw new Error('TFJS models need to be compiled to be used') @@ -176,12 +179,14 @@ export class TFJS extends Model { static async deserialize([ datatype, artifacts, + metadata ]: Serialized): Promise> { return new this( datatype, await tf.loadLayersModel({ load: () => Promise.resolve(artifacts), }), + metadata ); } @@ -204,7 +209,7 @@ export class TFJS extends Model { includeOptimizer: true // keep model compiled }) - return [this.datatype, await ret] + return [this.datatype, await ret, this.metadata] } [Symbol.dispose](): void{ diff --git a/discojs/src/processing/index.ts b/discojs/src/processing/index.ts index 4011f40d1..44ad656c2 100644 --- a/discojs/src/processing/index.ts +++ b/discojs/src/processing/index.ts @@ -9,6 +9,7 @@ import type { Tabular, Task, Network, + ModelMetadata, } from "../index.js"; import * as processing from "./index.js"; @@ -19,6 +20,7 @@ export * from "./tabular.js"; export function preprocess( task: Task, dataset: Dataset, + metadata?: ModelMetadata, ): Dataset { switch (task.dataType) { case "image": { @@ -37,12 +39,17 @@ export function preprocess( // cast as typescript doesn't reduce generic type const d = dataset as Dataset; const { inputColumns, outputColumn } = task.trainingInformation; + const stats = metadata?.tabularStandardization; return d.map((row) => { const output = processing.extractColumn(row, outputColumn); + const inputs = stats + ? List(processing.standardizeRow(row, inputColumns, stats)) + : extractToNumbers(inputColumns, row); + return [ - extractToNumbers(inputColumns, row), + inputs, // TODO sanitization doesn't care about column distribution output !== "" ? processing.convertToNumber(output) : 0, ]; @@ -68,6 +75,7 @@ export function preprocess( export function preprocessWithoutLabel( task: Task, dataset: Dataset, + metadata?: ModelMetadata, ): Dataset { switch (task.dataType) { case "image": { @@ -85,8 +93,13 @@ export function preprocessWithoutLabel( // cast as typescript doesn't reduce generic type const d = dataset as Dataset; const { inputColumns } = task.trainingInformation; + const stats = metadata?.tabularStandardization; - return d.map((row) => extractToNumbers(inputColumns, row)); + return d.map((row) => + stats + ? List(processing.standardizeRow(row, inputColumns, stats)) + : extractToNumbers(inputColumns, row) + ); } case "text": { // cast as typescript doesn't reduce generic type diff --git a/discojs/src/processing/tabular.ts b/discojs/src/processing/tabular.ts index 685baca03..73243428b 100644 --- a/discojs/src/processing/tabular.ts +++ b/discojs/src/processing/tabular.ts @@ -1,5 +1,10 @@ import { List } from "immutable"; +export type StandardizationStats = { + means: Record; + stds: Record; +}; + /** * Convert a string to a number * @@ -38,3 +43,57 @@ export function indexInList( if (ret === -1) throw new Error(`${element} not found in list`); return ret; } + +/** + * Return the mean, std value of each column + */ +export function computeStandardizationStats( + rows: Array>>, + columns: Array, +): StandardizationStats{ + const means: Record = {}; + const stds: Record = {}; + + for (const col of columns){ + const values = rows.map((row)=> convertToNumber(extractColumn(row, col))); + const mean = values.reduce((a, b)=> a+b, 0) / values.length; + const variance = values.reduce((acc, val) => acc + (val-mean)**2, 0) / values.length; + + const std = Math.sqrt(variance); + + means[col] = mean; + stds[col] = std; + } + + return {means, stds}; +} + +/** + * Apply standardization for a single value + */ +export function standardizeValue( + value: number, + mean: number, + std: number, +): number{ + if (std == 0) return 0; // avoid divide by 0 + return (value - mean) / std; +} + +/** + * Apply standardization for a row + * + * standardization function is called for each row in dataset + */ +export function standardizeRow( + row: Partial>, + columns: Array, + stats: StandardizationStats, +): Array{ + return columns.map((col) => { + const value = convertToNumber(extractColumn(row, col)); + const mean = stats.means[col]; + const std = stats.stds[col]; + return standardizeValue(value, mean, std); + }) +} \ No newline at end of file diff --git a/discojs/src/serialization/model.ts b/discojs/src/serialization/model.ts index 020d147af..b04f50f76 100644 --- a/discojs/src/serialization/model.ts +++ b/discojs/src/serialization/model.ts @@ -1,6 +1,6 @@ import type tf from '@tensorflow/tfjs' -import type { DataType, Model } from '../index.js' +import type { DataType, Model, ModelMetadata } from '../index.js' import { models, serialization } from '../index.js' import { GPTConfig } from '../models/index.js' @@ -41,11 +41,11 @@ export async function decode(encoded: Encoded): Promise> { const rawModel = raw[1] as unknown switch (type) { case Type.TFJS: { - if (raw.length !== 3) + if (raw.length !== 3 && raw.length !== 4) throw new Error( - "invalid TFJS model encoding: should be an array of length 3", + "invalid TFJS model encoding: should be an array of length 3 or 4", ); - const [rawDatatype, rawModel] = raw.slice(1) as unknown[]; + const [rawDatatype, rawModel, rawMetadata] = raw.slice(1) as unknown[]; let datatype; switch (rawDatatype) { @@ -63,6 +63,8 @@ export async function decode(encoded: Encoded): Promise> { datatype, // TODO totally unsafe casting rawModel as tf.io.ModelArtifacts, + // metadata for tabular task standardization + rawMetadata as ModelMetadata, ]); } case Type.GPT: { diff --git a/discojs/src/training/disco.ts b/discojs/src/training/disco.ts index 0d182fc43..f0b66fb54 100644 --- a/discojs/src/training/disco.ts +++ b/discojs/src/training/disco.ts @@ -155,13 +155,13 @@ export class Disco extends EventEmitter<{ > { this.#logger.success("Training started"); - const [trainingDataset, validationDataset] = - await this.#preprocessSplitAndBatch(dataset); - // the client fetches the latest weights upon connection // TODO unsafe cast this.trainer.model = (await this.#client.connect()) as Model; + const [trainingDataset, validationDataset] = + await this.#preprocessSplitAndBatch(dataset); + for await (const [round, epochs] of enumerate( this.trainer.train(trainingDataset, validationDataset), )) { @@ -213,21 +213,78 @@ export class Disco extends EventEmitter<{ > { const { batchSize, validationSplit } = this.#task.trainingInformation; - let preprocessed = processing.preprocess(this.#task, dataset); + if (validationSplit === 0){ + if (this.#task.dataType === "tabular"){ + const rows = await arrayFromAsync(dataset as Dataset); + const inputColumns = this.#task.trainingInformation.inputColumns; + + const stats = processing.computeStandardizationStats(rows, inputColumns); + this.trainer.model.metadata = { + tabularStandardization: stats, + }; + + let preprocessed = processing.preprocess( + this.#task, + dataset, + this.trainer.model.metadata, + ); + return [preprocessed.batch(batchSize).cached(), undefined]; + } + // If task datatype is not tabular + let preprocessed = processing.preprocess(this.#task, dataset); + + preprocessed = ( + this.#preprocessOnce + ? new Dataset(await arrayFromAsync(preprocessed)) + : preprocessed + ) + return [preprocessed.batch(batchSize).cached(), undefined]; + } + + // If training/validation splitting ratio is defined + const [training, validation] = dataset.split(validationSplit); + + if (this.#task.dataType == "tabular"){ + const trainingRows = await arrayFromAsync(training as Dataset); + const inputColumns = this.#task.trainingInformation.inputColumns; + const stats = processing.computeStandardizationStats(trainingRows, inputColumns); + + this.trainer.model.metadata = { + tabularStandardization: stats, + }; + + let preprocessedTraining = processing.preprocess(this.#task, training, this.trainer.model.metadata); + let preprocessedValidation = processing.preprocess(this.#task, validation, this.trainer.model.metadata); + preprocessedTraining = this.#preprocessOnce + ? new Dataset(await arrayFromAsync(preprocessedTraining)) + : preprocessedTraining; + + preprocessedValidation = this.#preprocessOnce + ? new Dataset(await arrayFromAsync(preprocessedValidation)) + : preprocessedValidation; + + return [ + preprocessedTraining.batch(batchSize).cached(), + preprocessedValidation.batch(batchSize).cached(), + ]; + } + + // if task datatype is not tabular + let preprocessedTraining = processing.preprocess(this.#task, training); + let preprocessedValidation = processing.preprocess(this.#task, validation); - preprocessed = ( - this.#preprocessOnce - ? new Dataset(await arrayFromAsync(preprocessed)) - : preprocessed - ) - if (validationSplit === 0) return [preprocessed.batch(batchSize).cached(), undefined]; + preprocessedTraining = this.#preprocessOnce + ? new Dataset(await arrayFromAsync(preprocessedTraining)) + : preprocessedTraining; - const [training, validation] = preprocessed.split(validationSplit); + preprocessedValidation = this.#preprocessOnce + ? new Dataset(await arrayFromAsync(preprocessedValidation)) + : preprocessedValidation; return [ - training.batch(batchSize).cached(), - validation.batch(batchSize).cached(), - ]; + preprocessedTraining.batch(batchSize).cached(), + preprocessedValidation.batch(batchSize).cached(), + ]; } } From 32e02cb1b082e23bec4b55640e5bb677d54e19fc Mon Sep 17 00:00:00 2001 From: ahzero7d1 Date: Mon, 30 Mar 2026 18:02:57 +0200 Subject: [PATCH 3/6] Add empty string processing in standardization --- discojs/src/processing/tabular.ts | 10 ++++++++-- discojs/src/training/disco.ts | 2 +- 2 files changed, 9 insertions(+), 3 deletions(-) diff --git a/discojs/src/processing/tabular.ts b/discojs/src/processing/tabular.ts index 73243428b..79ffa96a5 100644 --- a/discojs/src/processing/tabular.ts +++ b/discojs/src/processing/tabular.ts @@ -55,7 +55,10 @@ export function computeStandardizationStats( const stds: Record = {}; for (const col of columns){ - const values = rows.map((row)=> convertToNumber(extractColumn(row, col))); + const values = rows.map((row)=> { + const rawValue = extractColumn(row, col); + return convertToNumber(rawValue !== "" ? rawValue : "0"); + }); const mean = values.reduce((a, b)=> a+b, 0) / values.length; const variance = values.reduce((acc, val) => acc + (val-mean)**2, 0) / values.length; @@ -91,7 +94,10 @@ export function standardizeRow( stats: StandardizationStats, ): Array{ return columns.map((col) => { - const value = convertToNumber(extractColumn(row, col)); + const rawValue = extractColumn(row, col) + // Handle cases where the dataset contains empty strings. + // This only occurs in test cases, as empty strings are not allowed in the web app. + const value = convertToNumber(rawValue !== "" ? rawValue : "0"); const mean = stats.means[col]; const std = stats.stds[col]; return standardizeValue(value, mean, std); diff --git a/discojs/src/training/disco.ts b/discojs/src/training/disco.ts index f0b66fb54..7f808ba83 100644 --- a/discojs/src/training/disco.ts +++ b/discojs/src/training/disco.ts @@ -223,7 +223,7 @@ export class Disco extends EventEmitter<{ tabularStandardization: stats, }; - let preprocessed = processing.preprocess( + const preprocessed = processing.preprocess( this.#task, dataset, this.trainer.model.metadata, From ced7768ce80690076406fa30dca17c0a9319c4ab Mon Sep 17 00:00:00 2001 From: ahzero7d1 Date: Fri, 11 Sep 2026 09:54:05 +0200 Subject: [PATCH 4/6] Add one-hot encoding & categorical feature information to training information --- discojs/src/default_tasks/titanic.ts | 1 + discojs/src/processing/index.spec.ts | 1 + discojs/src/processing/index.ts | 16 ++-- discojs/src/processing/tabular.ts | 71 +++++++++++--- discojs/src/task/training_information.ts | 2 + discojs/src/training/disco.ts | 13 ++- webapp/cypress/e2e/datasetInput.cy.ts | 1 + webapp/cypress/e2e/task-creation.cy.ts | 1 + .../task_creation_form/TaskCreationForm.vue | 94 ++++++++++++++++++- 9 files changed, 172 insertions(+), 28 deletions(-) diff --git a/discojs/src/default_tasks/titanic.ts b/discojs/src/default_tasks/titanic.ts index 5ec8829d7..4b5d53340 100644 --- a/discojs/src/default_tasks/titanic.ts +++ b/discojs/src/default_tasks/titanic.ts @@ -48,6 +48,7 @@ export const titanic: TaskProvider<"tabular", "federated"> = { 'Fare', 'Pclass' ], + categoricalColumns: {}, outputColumn: 'Survived', scheme: 'federated', aggregationStrategy: 'mean', diff --git a/discojs/src/processing/index.spec.ts b/discojs/src/processing/index.spec.ts index 19aa2b47a..58059584d 100644 --- a/discojs/src/processing/index.spec.ts +++ b/discojs/src/processing/index.spec.ts @@ -23,6 +23,7 @@ describe("preprocess", () => { batchSize: 1, validationSplit: 0, inputColumns: ["a", "b"], + categoricalColumns: {}, outputColumn: "c", }, }; diff --git a/discojs/src/processing/index.ts b/discojs/src/processing/index.ts index 44ad656c2..5582c46bc 100644 --- a/discojs/src/processing/index.ts +++ b/discojs/src/processing/index.ts @@ -38,15 +38,15 @@ export function preprocess( case "tabular": { // cast as typescript doesn't reduce generic type const d = dataset as Dataset; - const { inputColumns, outputColumn } = task.trainingInformation; + const { inputColumns, outputColumn, categoricalColumns } = task.trainingInformation; const stats = metadata?.tabularStandardization; return d.map((row) => { const output = processing.extractColumn(row, outputColumn); - const inputs = stats - ? List(processing.standardizeRow(row, inputColumns, stats)) - : extractToNumbers(inputColumns, row); + const inputs = List( + processing.encodeTabularRow(row, inputColumns, categoricalColumns, stats) + ); return [ inputs, @@ -92,13 +92,13 @@ export function preprocessWithoutLabel( case "tabular": { // cast as typescript doesn't reduce generic type const d = dataset as Dataset; - const { inputColumns } = task.trainingInformation; + const { inputColumns, categoricalColumns } = task.trainingInformation; const stats = metadata?.tabularStandardization; return d.map((row) => - stats - ? List(processing.standardizeRow(row, inputColumns, stats)) - : extractToNumbers(inputColumns, row) + List( + processing.encodeTabularRow(row, inputColumns, categoricalColumns, stats) + ) ); } case "text": { diff --git a/discojs/src/processing/tabular.ts b/discojs/src/processing/tabular.ts index 79ffa96a5..1446130cf 100644 --- a/discojs/src/processing/tabular.ts +++ b/discojs/src/processing/tabular.ts @@ -84,22 +84,63 @@ export function standardizeValue( } /** - * Apply standardization for a row + * Apply one hot encoding for a row * - * standardization function is called for each row in dataset + * One hot encoding function is called for each row in dataset */ -export function standardizeRow( +export function oneHotEncode( + value: string, + categories: Array, +): Array { + // Get the index of the value among the possible categories + const index = categories.indexOf(value); + + // If the value does not exist, raise an error + if (index === -1) { + throw new Error(`"${value}" is not a valid category for this column`); + } + + return categories.map((_, categoryIndex) => + categoryIndex === index ? 1 : 0 + ); +} + +/** + * Apply standardization for numerical columns and + * apply one hot encoding for categorical columns and return the final row + */ +export function encodeTabularRow( row: Partial>, - columns: Array, - stats: StandardizationStats, -): Array{ - return columns.map((col) => { - const rawValue = extractColumn(row, col) - // Handle cases where the dataset contains empty strings. - // This only occurs in test cases, as empty strings are not allowed in the web app. - const value = convertToNumber(rawValue !== "" ? rawValue : "0"); - const mean = stats.means[col]; - const std = stats.stds[col]; - return standardizeValue(value, mean, std); - }) + inputColumns: Array, + categoricalColumns: Record>, + stats?: StandardizationStats, +): Array { + const outputRow = inputColumns.flatMap((column) => { + const raw = extractColumn(row, column); + const categories = categoricalColumns[column]; + + // If the column exists in the list of categorical columns, apply one hot encoding + if (categories !== undefined){ + return oneHotEncode(raw, categories); + } + + // If the column is numerical column, apply standardization + const value = convertToNumber(raw !== "" ? raw : "0"); + + if (stats === undefined) { + return [value]; + } + + const mean = stats.means[column]; + const std = stats.stds[column]; + + // Raise an error when stats is not defined + if (mean === undefined || std === undefined){ + throw new Error(`Standardization statistics is not defined for column ${column}`); + } + + return [standardizeValue(value, mean, std)]; + }); + + return outputRow; } \ No newline at end of file diff --git a/discojs/src/task/training_information.ts b/discojs/src/task/training_information.ts index fefb41f40..745e9c9f4 100644 --- a/discojs/src/task/training_information.ts +++ b/discojs/src/task/training_information.ts @@ -83,6 +83,8 @@ export namespace TrainingInformation { tabular: z.object({ // the columns to be chosen as input data for the model inputColumns: z.array(z.string()), + // categorical columns to be chosen as input data for the model + categoricalColumns: z.record(z.string(), z.array(z.string()).min(1)).optional().default({}), // the columns to be predicted by the model outputColumn: z.string(), }), diff --git a/discojs/src/training/disco.ts b/discojs/src/training/disco.ts index 7f808ba83..ac4bc5fa3 100644 --- a/discojs/src/training/disco.ts +++ b/discojs/src/training/disco.ts @@ -218,7 +218,11 @@ export class Disco extends EventEmitter<{ const rows = await arrayFromAsync(dataset as Dataset); const inputColumns = this.#task.trainingInformation.inputColumns; - const stats = processing.computeStandardizationStats(rows, inputColumns); + // Make sure to compute standardization stats for numerical features + const categoricalColumns = new Set(Object.keys(this.#task.trainingInformation.categoricalColumns)); + const numericalColumns = inputColumns.filter(column => !categoricalColumns.has(column)); + + const stats = processing.computeStandardizationStats(rows, numericalColumns); this.trainer.model.metadata = { tabularStandardization: stats, }; @@ -247,7 +251,12 @@ export class Disco extends EventEmitter<{ if (this.#task.dataType == "tabular"){ const trainingRows = await arrayFromAsync(training as Dataset); const inputColumns = this.#task.trainingInformation.inputColumns; - const stats = processing.computeStandardizationStats(trainingRows, inputColumns); + + // Make sure to compute standardization stats for numerical features + const categoricalColumns = new Set(Object.keys(this.#task.trainingInformation.categoricalColumns)); + const numericalColumns = inputColumns.filter(column => !categoricalColumns.has(column)); + + const stats = processing.computeStandardizationStats(trainingRows, numericalColumns); this.trainer.model.metadata = { tabularStandardization: stats, diff --git a/webapp/cypress/e2e/datasetInput.cy.ts b/webapp/cypress/e2e/datasetInput.cy.ts index d26a83eed..2c512cf63 100644 --- a/webapp/cypress/e2e/datasetInput.cy.ts +++ b/webapp/cypress/e2e/datasetInput.cy.ts @@ -86,6 +86,7 @@ describe("tabular dataset input", () => { setupServerWith( basicTask("tabular", { inputColumns: ["a", "b"], + categoricalColumns: {}, outputColumn: "c", }), ); diff --git a/webapp/cypress/e2e/task-creation.cy.ts b/webapp/cypress/e2e/task-creation.cy.ts index 4ea2c9ecf..679866e4e 100644 --- a/webapp/cypress/e2e/task-creation.cy.ts +++ b/webapp/cypress/e2e/task-creation.cy.ts @@ -85,6 +85,7 @@ it("submits with tabular task", () => { validationSplit: 0, minNbOfParticipants: 2, inputColumns: ["input"], + categoricalColumns: {}, outputColumn: "output", tensorBackend: "tfjs", }, diff --git a/webapp/src/components/task_creation_form/TaskCreationForm.vue b/webapp/src/components/task_creation_form/TaskCreationForm.vue index 5ff1d705a..9f126d72a 100644 --- a/webapp/src/components/task_creation_form/TaskCreationForm.vue +++ b/webapp/src/components/task_creation_form/TaskCreationForm.vue @@ -148,6 +148,84 @@ + + +
+
+ +
+ + + + + +
+ + + +
+
+ + + + + +
+ + + add category + +
+
+
+ + + add categorical column + +
+
+
+ { break; case "tabular": form.setFieldValue("trainingInformation.inputColumns", [""]); + form.setFieldValue("trainingInformation.categoricalColumns", []); break; } }); @@ -855,9 +934,18 @@ const schema = z z.object({ ...Task.dataTypeToSchema.tabular.shape, ...TFJSModelSchema, - trainingInformation: TrainingInformation.dataTypeToSchema.tabular.and( - trainingInformationNetworks, - ), + trainingInformation: TrainingInformation.dataTypeToSchema.tabular.extend({ + categoricalColumns: z.array( + z.object({ + column: z.string().trim().min(1), + categories: z.array(z.string().min(1)).min(1) + }) + ).default([]).transform((columns) => Object.fromEntries(columns.map(({column, categories}) => [ + column, + categories, + ]))) + }) + .and(trainingInformationNetworks,), }), z.object({ ...Task.dataTypeToSchema.text.shape, From a6f86f395889494cea72893134db9b25150e9c9a Mon Sep 17 00:00:00 2001 From: ahzero7d1 Date: Fri, 11 Sep 2026 13:45:19 +0200 Subject: [PATCH 5/6] Improve input form & check if categorical columns are included in input columns --- .../task_creation_form/TaskCreationForm.vue | 45 ++++++++++++++----- 1 file changed, 33 insertions(+), 12 deletions(-) diff --git a/webapp/src/components/task_creation_form/TaskCreationForm.vue b/webapp/src/components/task_creation_form/TaskCreationForm.vue index 9f126d72a..ba25bbab4 100644 --- a/webapp/src/components/task_creation_form/TaskCreationForm.vue +++ b/webapp/src/components/task_creation_form/TaskCreationForm.vue @@ -134,7 +134,7 @@ > @@ -167,14 +167,20 @@ class="flex flex-col elems-gap" > -
- - - +
+
+ +
+ +
@@ -188,7 +194,7 @@ }" :name="`trainingInformation.categoricalColumns[${i}].categories`" > -
+
@@ -945,7 +952,7 @@ const schema = z categories, ]))) }) - .and(trainingInformationNetworks,), + .and(trainingInformationNetworks), }), z.object({ ...Task.dataTypeToSchema.text.shape, @@ -968,7 +975,21 @@ const schema = z .and(trainingInformationNetworks), }), ]), - ); + ).superRefine((task, ctx) => { + if (task.dataType !== "tabular") return; + + const {inputColumns, categoricalColumns} = task.trainingInformation; + + Object.keys(categoricalColumns).forEach((column, idx) => { + if (inputColumns.includes(column)) return; + + ctx.addIssue({ + code: "custom", + path: ["trainingInformation", "categoricalColumns", idx, "column"], + message: "Categorical columns must also be included in input columns" + }); + }); + }); async function onSubmit(form: unknown): Promise { // TODO double check as @submit isn't generic vee-validate#4845 From 16c6393fb56c1ddbc9ad80a1cdbb387ad6609338 Mon Sep 17 00:00:00 2001 From: ahzero7d1 Date: Fri, 11 Sep 2026 15:53:36 +0200 Subject: [PATCH 6/6] Update titanic default task --- discojs/src/default_tasks/titanic.ts | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/discojs/src/default_tasks/titanic.ts b/discojs/src/default_tasks/titanic.ts index 4b5d53340..a2b6e3e7a 100644 --- a/discojs/src/default_tasks/titanic.ts +++ b/discojs/src/default_tasks/titanic.ts @@ -46,9 +46,14 @@ export const titanic: TaskProvider<"tabular", "federated"> = { 'SibSp', 'Parch', 'Fare', - 'Pclass' + 'Pclass', + 'Sex', + 'Embarked' ], - categoricalColumns: {}, + categoricalColumns: { + Sex: ["male", "female"], + Embarked: ["C", "S", "Q", "Missing"] + }, outputColumn: 'Survived', scheme: 'federated', aggregationStrategy: 'mean', @@ -63,7 +68,7 @@ export const titanic: TaskProvider<"tabular", "federated"> = { model.add( tf.layers.dense({ - inputShape: [5], + inputShape: [11], units: 124, activation: 'relu', kernelInitializer: 'leCunNormal'