@astermind/astermind-pro
Version:
Astermind Pro - Premium ML Toolkit with Advanced RAG, Reranking, Summarization, and Information Flow Analysis
158 lines • 6.01 kB
JavaScript
// 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