Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions .changeset/browser-optimization-runtime.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
"@hashintel/petrinaut-core": patch
---

Adds an in-browser optimization capability that runs the Optuna study in a Pyodide worker and evaluates trials through a host channel. A study that completed or was cancelled stays in the worker until it is released, so the connected capability can extend it with more trials on the same sampler history, and a run may keep up to four trials in flight at once.
159 changes: 15 additions & 144 deletions libs/@hashintel/petrinaut-cli/src/runtime/optimization.ts
Original file line number Diff line number Diff line change
Expand Up @@ -3,32 +3,28 @@ import { readFile } from "node:fs/promises";
import {
compileScenario,
createMonteCarloExperiment,
deriveRunSeed,
parseDocumentText,
petrinautOptimizationEvaluateParamsSchema,
petrinautOptimizationManifestSchema,
} from "@hashintel/petrinaut-core";
import { lowerScenarioToHir } from "@hashintel/petrinaut-core/hir";
import {
deriveOptimizationTrialSeeds,
describeOptimization,
resolveTrialScenarioParameterValues,
} from "@hashintel/petrinaut-core/optimization";
import { createInProcessMonteCarloWorker } from "@hashintel/petrinaut-core/workers/monte-carlo";

import type {
MonteCarloExperiment,
PetrinautOptimizationDescribeParameter,
PetrinautOptimizationDescribeResult,
PetrinautOptimizationEvaluateResult,
PetrinautOptimizationManifest,
Scenario,
WorkerFactory,
} from "@hashintel/petrinaut-core";
import type { PetrinautCompiledModel } from "@hashintel/petrinaut-core/compiled-model";

type OptimizationScalar = number | boolean;
type ScenarioParameter = Scenario["scenarioParameters"][number];
type OptimizedBinding = Extract<
PetrinautOptimizationManifest["scenario"]["parameterBindings"][string],
{ kind: "optimize" }
>;
type OptimizationDomain = OptimizedBinding["domain"];
export { deriveOptimizationTrialSeeds as deriveTrialSeeds } from "@hashintel/petrinaut-core/optimization";

function formatManifestIssues(
prefix: string,
Expand Down Expand Up @@ -67,95 +63,11 @@ export async function loadOptimizationManifest(
return parseOptimizationManifest(document.data);
}

function describeParameter(
parameter: ScenarioParameter,
domain: OptimizationDomain,
): PetrinautOptimizationDescribeParameter {
switch (domain.kind) {
case "continuous":
return {
identifier: parameter.identifier,
type: "float",
default: parameter.default,
minimum: domain.minimum,
maximum: domain.maximum,
scale: domain.scale,
};
case "integer":
return {
identifier: parameter.identifier,
type: "int",
default: parameter.default,
minimum: domain.minimum,
maximum: domain.maximum,
step: domain.step,
scale: domain.scale,
};
case "boolean":
return {
identifier: parameter.identifier,
type: "boolean",
default: parameter.default !== 0,
};
}
}

function validateSuggestedValue(
parameter: ScenarioParameter,
domain: OptimizationDomain,
value: OptimizationScalar,
): void {
if (domain.kind === "boolean") {
if (typeof value !== "boolean") {
throw new Error(
`Optimization parameter "${parameter.identifier}" must be boolean`,
);
}
return;
}
if (typeof value !== "number") {
throw new Error(
`Optimization parameter "${parameter.identifier}" must be numeric`,
);
}
if (value < domain.minimum || value > domain.maximum) {
throw new Error(
`Optimization parameter "${parameter.identifier}" must be between ${domain.minimum} and ${domain.maximum}`,
);
}
if (domain.kind === "integer") {
if (!Number.isInteger(value)) {
throw new Error(
`Optimization parameter "${parameter.identifier}" must be an integer`,
);
}
if ((value - domain.minimum) % domain.step !== 0) {
throw new Error(
`Optimization parameter "${parameter.identifier}" must align with step ${domain.step} from ${domain.minimum}`,
);
}
}
}

export type OptimizationProtocol = {
describe(): PetrinautOptimizationDescribeResult;
evaluate(params: unknown): Promise<PetrinautOptimizationEvaluateResult>;
};

/**
* Derives one trial's run seeds. Run 0 keeps the base seed, so a single-seed
* trial matches the old fixed-seed behaviour; later runs use the Monte Carlo
* derivation. Every trial gets the same sequence: common random numbers.
*/
export function deriveTrialSeeds(
baseSeed: number,
seedsPerTrial: number,
): number[] {
return Array.from({ length: seedsPerTrial }, (_, index) =>
index === 0 ? baseSeed : deriveRunSeed(baseSeed, index),
);
}

/** Resolves when the experiment reports its terminal event. */
function waitForCompletion(experiment: MonteCarloExperiment): Promise<void> {
return new Promise((resolve, reject) => {
Expand Down Expand Up @@ -204,25 +116,17 @@ export function createOptimizationProtocol(args: {
const createWorker = args.createWorker ?? createInProcessMonteCarloWorker;
const createExperiment = args.createExperiment ?? createMonteCarloExperiment;
const seedsPerTrial = manifest.execution.seedsPerTrial ?? 1;
const trialSeeds = deriveTrialSeeds(manifest.execution.seed, seedsPerTrial);
const trialSeeds = deriveOptimizationTrialSeeds(
manifest.execution.seed,
seedsPerTrial,
);
const scenario = manifest.model.definition.scenarios?.[0];
const metric = manifest.model.definition.metrics?.[0];
if (!scenario || !metric) {
throw new Error(
"An optimization manifest requires exactly one scenario and one metric",
);
}
const optimizedParameters = scenario.scenarioParameters.flatMap(
(parameter) => {
const binding = manifest.scenario.parameterBindings[parameter.identifier];
return binding?.kind === "optimize"
? [{ parameter, domain: binding.domain }]
: [];
},
);
const optimizedIdentifiers = new Set(
optimizedParameters.map(({ parameter }) => parameter.identifier),
);

// Lower the scenario's expressions once per study; each trial re-runs only
// the type-check and the interpreter with that trial's parameter values.
Expand All @@ -237,17 +141,7 @@ export function createOptimizationProtocol(args: {

return {
describe() {
return {
direction: manifest.objective.direction,
study: {
...manifest.study,
seed: manifest.execution.seed,
seedsPerTrial,
},
parameters: optimizedParameters.map(({ parameter, domain }) =>
describeParameter(parameter, domain),
),
};
return describeOptimization(manifest);
},
async evaluate(params) {
const parsed =
Expand All @@ -258,33 +152,10 @@ export function createOptimizationProtocol(args: {
parsed.error.issues,
);
}
const values = parsed.data.parameterValues;
for (const { parameter } of optimizedParameters) {
const { identifier } = parameter;
if (!Object.hasOwn(values, identifier)) {
throw new Error(`Missing optimized parameter "${identifier}"`);
}
}
for (const identifier of Object.keys(values)) {
if (!optimizedIdentifiers.has(identifier)) {
throw new Error(`Unexpected optimization parameter "${identifier}"`);
}
}

const scenarioParameterValues: Record<string, number> = {};
for (const parameter of scenario.scenarioParameters) {
const binding =
manifest.scenario.parameterBindings[parameter.identifier]!;
const value =
binding.kind === "fixed"
? binding.value
: values[parameter.identifier]!;
if (binding.kind === "optimize") {
validateSuggestedValue(parameter, binding.domain, value);
}
scenarioParameterValues[parameter.identifier] =
typeof value === "boolean" ? (value ? 1 : 0) : value;
}
const scenarioParameterValues = resolveTrialScenarioParameterValues(
manifest,
parsed.data.parameterValues,
);

const compiledScenario = compileScenario(
scenario,
Expand Down
6 changes: 5 additions & 1 deletion libs/@hashintel/petrinaut-core/.oxlintrc.json
Original file line number Diff line number Diff line change
Expand Up @@ -89,7 +89,11 @@
{
"patterns": [
{
"group": ["@local/*"],
"group": [
"@local/*",
"!@local/petrinaut-optimizer-core",
"!@local/petrinaut-optimizer-core/**"
],
"message": "You cannot use unpublished local packages in a published package."
},
{
Expand Down
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
{
"package": "@hashintel/petrinaut-core",
"dependencies": [],
"dependencies": [
"@local/petrinaut-optimizer-core"
],
"tasks": {
"build": [],
"fix:eslint": [
Expand Down
10 changes: 10 additions & 0 deletions libs/@hashintel/petrinaut-core/package.json
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,10 @@
"types": "./dist/ai.d.ts",
"import": "./dist/ai.js"
},
"./browser-optimization": {
"types": "./dist/browser-optimization.d.ts",
"import": "./dist/browser-optimization.js"
},
"./compiled-model": {
"types": "./dist/compiled-model.d.ts",
"import": "./dist/compiled-model.js"
Expand Down Expand Up @@ -66,6 +70,10 @@
"types": "./dist/workers/monte-carlo.d.ts",
"import": "./dist/workers/monte-carlo.js"
},
"./workers/optimizer": {
"types": "./dist/workers/optimizer.d.ts",
"import": "./dist/workers/optimizer.js"
},
"./workers/simulation": {
"types": "./dist/workers/simulation.d.ts",
"import": "./dist/workers/simulation.js"
Expand Down Expand Up @@ -98,12 +106,14 @@
"zod": "4.4.3"
},
"devDependencies": {
"@local/petrinaut-optimizer-core": "workspace:*",
"@types/js-yaml": "^4",
"@types/node": "22.18.13",
"@typescript/native-preview": "7.0.0-dev.20260511.1",
"@webgpu/types": "0.1.71",
"oxlint": "1.63.0",
"oxlint-tsgolint": "0.22.1",
"pyodide": "314.0.6",
"rolldown": "1.2.6",
"rolldown-plugin-dts": "0.28.3",
"typescript": "5.9.3",
Expand Down
22 changes: 22 additions & 0 deletions libs/@hashintel/petrinaut-core/src/browser-optimization.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
export {
createBrowserOptimization,
type CreateBrowserOptimizationOptions,
} from "./browser-optimization/browser-optimization";
export type {
OptimizerWorkerErrorEvent,
OptimizerWorkerLike,
} from "./browser-optimization/create-optimizer-worker";
export {
defaultOptimizerPyodideConfig,
type OptimizerPyodideConfig,
} from "./browser-optimization/pyodide-config";
export type {
OptimizationScalar,
PetrinautConnectedOptimization,
PetrinautConnectedOptimizationCapability,
PetrinautConnectedRunOptions,
PetrinautOptimizationChannel,
PetrinautOptimizationSource,
PetrinautOptimizationTrialOutcome,
PetrinautOptimizationTrialRequest,
} from "./optimization";
Loading
Loading