UNPKG

@astermind/astermind-pro

Version:

Astermind Pro - Premium ML Toolkit with Advanced RAG, Reranking, Summarization, and Information Flow Analysis

158 lines 6.01 kB
// multi-task-elm.ts — Multi-Task ELM for joint learning across related tasks // Shared hidden layer with task-specific output layers import { ELM } from '@astermind/astermind-elm'; import { requireLicense } from '../core/license.js'; /** * Multi-Task ELM for joint learning across related tasks * Features: * - Shared feature extraction layer * - Task-specific output layers * - Task weighting for importance * - Joint optimization */ export class MultiTaskELM { constructor(options) { this.taskELMs = new Map(); this.trained = false; requireLicense(); // Premium feature - requires valid license this.tasks = options.tasks.map((task) => ({ name: task.name, categories: task.categories, weight: task.weight ?? 1.0, })); this.options = { sharedHiddenUnits: options.sharedHiddenUnits ?? 256, taskSpecificHiddenUnits: options.taskSpecificHiddenUnits ?? options.tasks.map(() => 128), activation: options.activation ?? 'relu', maxLen: options.maxLen ?? 100, useTokenizer: options.useTokenizer ?? true, }; // Initialize shared ELM this.sharedELM = new ELM({ useTokenizer: this.options.useTokenizer ? true : undefined, hiddenUnits: this.options.sharedHiddenUnits, categories: [], // No categories for shared layer maxLen: this.options.maxLen, activation: this.options.activation, }); // Initialize task-specific ELMs for (let i = 0; i < this.tasks.length; i++) { const task = this.tasks[i]; const taskELM = new ELM({ hiddenUnits: this.options.taskSpecificHiddenUnits[i], categories: task.categories, maxLen: this.options.sharedHiddenUnits, // Input size is shared layer output activation: this.options.activation, }); this.taskELMs.set(task.name, taskELM); } } /** * Train multi-task ELM * @param X Input features * @param yTaskData Map of task name to labels */ train(X, yTaskData) { // Step 1: Train shared layer (use all tasks) const allFeatures = this._extractSharedFeatures(X); // Step 2: Train each task-specific layer for (const task of this.tasks) { const taskLabels = yTaskData.get(task.name); if (!taskLabels) continue; const taskELM = this.taskELMs.get(task.name); const labelIndices = taskLabels.map(label => typeof label === 'number' ? label : task.categories.indexOf(label)); // Train task-specific ELM on shared features taskELM.setCategories?.(task.categories); taskELM.trainFromData?.(allFeatures, labelIndices); } this.trained = true; } /** * Predict for all tasks */ predict(X, topK = 3) { if (!this.trained) { throw new Error('Model must be trained before prediction'); } const XArray = Array.isArray(X[0]) ? X : [X]; const results = new Map(); for (const x of XArray) { // Extract shared features const sharedFeatures = this._extractSharedFeatures([x])[0]; // Predict for each task for (const task of this.tasks) { const taskELM = this.taskELMs.get(task.name); const taskPreds = taskELM.predictFromVector?.([sharedFeatures], topK) || []; const taskResults = taskPreds.map((pred) => ({ task: task.name, label: pred.label || task.categories[pred.index || 0], prob: pred.prob || 0, })); if (!results.has(task.name)) { results.set(task.name, []); } results.get(task.name).push(...taskResults); } } return results; } /** * Predict for a specific task */ predictTask(x, taskName, topK = 3) { if (!this.trained) { throw new Error('Model must be trained before prediction'); } const taskELM = this.taskELMs.get(taskName); if (!taskELM) { throw new Error(`Task ${taskName} not found`); } const XArray = Array.isArray(x[0]) ? x : [x]; const results = []; for (const xi of XArray) { // Extract shared features const sharedFeatures = this._extractSharedFeatures([xi])[0]; // Predict with task-specific ELM const taskPreds = taskELM.predictFromVector?.([sharedFeatures], topK) || []; results.push(...taskPreds.map((pred) => ({ task: taskName, label: pred.label || this.tasks.find(t => t.name === taskName).categories[pred.index || 0], prob: pred.prob || 0, }))); } return results; } /** * Extract features from shared layer */ _extractSharedFeatures(X) { // Encode inputs if using tokenizer const encoded = this.options.useTokenizer ? X.map(x => { const enc = this.sharedELM.encoder?.encode?.(x) || x; return this.sharedELM.encoder?.normalize?.(enc) || enc; }) : X; // Extract hidden layer features return encoded.map(x => { const hidden = this.sharedELM.buildHidden?.([x], this.sharedELM.model?.W, this.sharedELM.model?.b); return hidden?.[0] ? Array.from(hidden[0]) : x; }); } /** * Get task names */ getTaskNames() { return this.tasks.map(t => t.name); } /** * Get task weights */ getTaskWeights() { return new Map(this.tasks.map(t => [t.name, t.weight])); } } //# sourceMappingURL=multi-task-elm.js.map