From 4918869976cebb49571e3517b6ddff1a61f0adaa Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Fri, 25 Sep 2026 14:56:49 +0200 Subject: [PATCH 01/23] fix: byzantine aggregator tensor leaks --- discojs/src/aggregator/byzantine.spec.ts | 47 +++++++++++++++ discojs/src/aggregator/byzantine.ts | 74 ++++++++++++------------ 2 files changed, 85 insertions(+), 36 deletions(-) diff --git a/discojs/src/aggregator/byzantine.spec.ts b/discojs/src/aggregator/byzantine.spec.ts index 404ceae39..d2e1971b7 100644 --- a/discojs/src/aggregator/byzantine.spec.ts +++ b/discojs/src/aggregator/byzantine.spec.ts @@ -1,6 +1,7 @@ import { Set } from "immutable"; import { describe, expect, it } from "vitest"; import fc from "fast-check"; +import * as tf from "@tensorflow/tfjs"; import { WeightsContainer } from "#weights/index"; import { ByzantineRobustAggregator } from "#aggregator/byzantine"; @@ -342,4 +343,50 @@ describe("ByzantineRobustAggregator", () => { expect(await run()).to.be.closeTo(await run(), 1e-6); }); + + it("keeps working after the caller disposes the previous aggregate", async () => { + const agg = new ByzantineRobustAggregator(0, 2, "absolute", 1.0, 2, 0.5); + agg.setNodes(Set(["a", "b"])); + + const p1 = agg.getPromiseForAggregation(); + agg.add("a", WeightsContainer.of([1]), 0); + agg.add("b", WeightsContainer.of([1]), 0); + // the caller owns the aggregate, disposing it must not break the next round + (await p1).dispose(); + + const p2 = agg.getPromiseForAggregation(); + agg.add("a", WeightsContainer.of([1]), 1); + agg.add("b", WeightsContainer.of([1]), 1); + const out = await p2; + // m = 0.5 * 1 + 0.5 * 0.5 + expect((await WSIntoArrays(out))[0][0]).to.be.closeTo(0.75, 1e-6); + out.dispose(); + agg.dispose(); + }); + + it("aggregation leaves no dangling tensors", async () => { + const baseline = tf.memory().numTensors; + + const agg = new ByzantineRobustAggregator(0, 3, "absolute", 1.0, 3, 0.5); + const ids = ["a", "b", "c"]; + agg.setNodes(Set(ids)); + + for (let round = 0; round < 3; round++) { + const contributions = ids.map((_, i) => + WeightsContainer.of([i, round], [100 * i]), + ); + const p = agg.getPromiseForAggregation(); + // update a contribution within the same round + const replaced = WeightsContainer.of([0, 0], [0]); + agg.add("a", replaced, round); + ids.forEach((id, i) => agg.add(id, contributions[i], round)); + (await p).dispose(); + replaced.dispose(); + contributions.forEach((c) => c.dispose()); + } + + // the aggregator only keeps the momentums and the previous aggregate + agg.dispose(); + expect(tf.memory().numTensors).to.equal(baseline); + }); }); diff --git a/discojs/src/aggregator/byzantine.ts b/discojs/src/aggregator/byzantine.ts index 3e60808bb..2b90c974e 100644 --- a/discojs/src/aggregator/byzantine.ts +++ b/discojs/src/aggregator/byzantine.ts @@ -1,9 +1,8 @@ import { Map } from "immutable"; import * as tf from "@tensorflow/tfjs"; -import type { WeightsContainer } from "#weights/index"; import type { NodeID } from "#client/types"; -import { avg } from "#weights/index"; +import { avg, WeightsContainer } from "#weights/index"; import { AggregationStep } from "#aggregator/aggregator"; import type { ThresholdType } from "#aggregator/multiround"; @@ -98,10 +97,16 @@ export class ByzantineRobustAggregator extends MultiRoundAggregator { nodeId, ); + // Replacing a contribution of the current round, free the previous one + const previous = this.contributions.getIn([0, nodeId]) as + | WeightsContainer + | undefined; + previous?.dispose(); + const prevMomentum = this.historyMomentums.get(nodeId); const newMomentum = prevMomentum ? contribution.mapWith(prevMomentum, (g, m) => - g.mul(1 - this.beta).add(m.mul(this.beta)), + tf.tidy(() => g.mul(1 - this.beta).add(m.mul(this.beta))), ) : contribution.map((g) => g.mul(1 - this.beta)); @@ -125,55 +130,52 @@ export class ByzantineRobustAggregator extends MultiRoundAggregator { } // Step 1: Initialize v using previous aggregate or mean of contributions - let v: WeightsContainer; - if (this.prevAggregate) { - v = this.prevAggregate.map((t) => tf.clone(t)); // Clone to avoid in-place modifications - } else { - v = avg(currentContributions.values()); - } - - const eps = tf.scalar(1e-12); - const one = tf.scalar(1); - const radius = tf.scalar(this.clippingRadius); + // Clone to avoid in-place modifications of the stored aggregate + let v = this.prevAggregate?.clone() ?? avg(currentContributions.values()); // Step 2: Iterative Centered Clipping for (let l = 0; l < this.maxIterations; l++) { + // Clip one contribution at a time so that the unclipped diff is freed + // right away rather than keeping all of them alive until the end const clippedDiffs = Array.from(currentContributions.values()).map( - (m) => { - const diff = m.sub(v); - - const norm = euclideanNorm(diff); - - const safeNorm = tf.maximum(norm, eps); - - const scale = tf.minimum(one, tf.div(radius, safeNorm)); - - const clipped = diff.mul(scale); - - norm.dispose(); - safeNorm.dispose(); - scale.dispose(); - - return clipped; - }, + (m) => + new WeightsContainer( + tf.tidy(() => { + const diff = m.sub(v); + const safeNorm = tf.maximum(euclideanNorm(diff), 1e-12); + const scale = tf.minimum( + 1, + tf.div(this.clippingRadius, safeNorm), + ); + return diff.mul(scale).weights; + }), + ), ); const avgClip = avg(clippedDiffs); - const newV = v.add(avgClip); - clippedDiffs.forEach((d) => d.dispose()); + const newV = v.add(avgClip); + avgClip.dispose(); - const oldV = v; + v.dispose(); v = newV; - oldV.dispose(); } - tf.dispose([eps, one, radius]); // Step 3: Update history - this.prevAggregate = v; + // Keep our own copy, the returned aggregate is owned (and disposed) by the caller + this.prevAggregate?.dispose(); + this.prevAggregate = v.clone(); return v; } + override dispose(): void { + this.historyMomentums.forEach((momentum) => momentum.dispose()); + this.historyMomentums = Map(); + this.prevAggregate?.dispose(); + this.prevAggregate = null; + super.dispose(); + } + override makePayloads( weights: WeightsContainer, ): Map { From f173f80ca6a03e9839cf8431faa06942bd0f38af Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Fri, 25 Sep 2026 16:05:26 +0200 Subject: [PATCH 02/23] fix: dispose aggregators on close --- discojs/src/training/disco.ts | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/discojs/src/training/disco.ts b/discojs/src/training/disco.ts index 097eb4905..278396f93 100644 --- a/discojs/src/training/disco.ts +++ b/discojs/src/training/disco.ts @@ -307,11 +307,15 @@ export class Disco extends EventEmitter<{ * Completely stops the ongoing training instance. */ async close(): Promise { - // Dispose the model tensor + // Dispose the model tensor and the aggregator's buffered tensors try { await this.#client.disconnect(); } finally { - this.trainer[Symbol.dispose](); + try { + this.trainer[Symbol.dispose](); + } finally { + this.#client.aggregator.dispose(); + } } } From f9301a3941341c4b6d08a7dbc76131af97a95e81 Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Fri, 25 Sep 2026 16:07:09 +0200 Subject: [PATCH 03/23] fix: percentile_clipping aggregator tensor leaks --- .../aggregator/percentile_clipping.spec.ts | 42 +++++++++++++++++++ discojs/src/aggregator/percentile_clipping.ts | 21 ++++++---- 2 files changed, 54 insertions(+), 9 deletions(-) diff --git a/discojs/src/aggregator/percentile_clipping.spec.ts b/discojs/src/aggregator/percentile_clipping.spec.ts index 39afa403b..be32661ce 100644 --- a/discojs/src/aggregator/percentile_clipping.spec.ts +++ b/discojs/src/aggregator/percentile_clipping.spec.ts @@ -1,5 +1,6 @@ import { Set } from "immutable"; import { describe, expect, it } from "vitest"; +import * as tf from "@tensorflow/tfjs"; import { WeightsContainer } from "#weights/index"; import { PercentileClippingAggregator } from "#aggregator/percentile_clipping"; @@ -170,4 +171,45 @@ describe("PercentileClippingAggregator", () => { expect(await run()).to.be.closeTo(await run(), 1e-6); }); + + it("keeps working after the caller disposes the previous aggregate", async () => { + const agg = new PercentileClippingAggregator(0, 2, "absolute", 0.5); + agg.setNodes(Set(["a", "b"])); + + const p1 = agg.getPromiseForAggregation(); + agg.add("a", WeightsContainer.of([1]), 0); + agg.add("b", WeightsContainer.of([1]), 0); + // the caller owns the aggregate, disposing it must not break the next round + (await p1).dispose(); + + const p2 = agg.getPromiseForAggregation(); + agg.add("a", WeightsContainer.of([2]), 1); + agg.add("b", WeightsContainer.of([2]), 1); + const out = await p2; + expect((await WSIntoArrays(out))[0][0]).to.be.closeTo(2, 1e-6); + out.dispose(); + agg.dispose(); + }); + + it("aggregation leaves no dangling tensors", async () => { + const baseline = tf.memory().numTensors; + + const agg = new PercentileClippingAggregator(0, 3, "absolute", 0.5); + const ids = ["a", "b", "c"]; + agg.setNodes(Set(ids)); + + for (let round = 0; round < 3; round++) { + const contributions = ids.map((_, i) => + WeightsContainer.of([i, round], [100 * i]), + ); + const p = agg.getPromiseForAggregation(); + ids.forEach((id, i) => agg.add(id, contributions[i], round)); + (await p).dispose(); + contributions.forEach((c) => c.dispose()); + } + + // the aggregator only keeps the previous aggregate + agg.dispose(); + expect(tf.memory().numTensors).to.equal(baseline); + }); }); diff --git a/discojs/src/aggregator/percentile_clipping.ts b/discojs/src/aggregator/percentile_clipping.ts index 7b564c0e3..cfa26e3a1 100644 --- a/discojs/src/aggregator/percentile_clipping.ts +++ b/discojs/src/aggregator/percentile_clipping.ts @@ -79,14 +79,9 @@ export class PercentileClippingAggregator extends MultiRoundAggregator { this.log(AggregationStep.AGGREGATE); // Step 1: Get the centering reference (previous aggregation or initial avg vector) - let centerReference: WeightsContainer; - if (this.prevAggregate) { - centerReference = this.prevAggregate.map((t) => tf.clone(t)); - } else { - centerReference = avg(currentContributions.values()).map((t) => - tf.clone(t), - ); - } + // Clone to avoid in-place modifications of the stored aggregate + const centerReference = + this.prevAggregate?.clone() ?? avg(currentContributions.values()); // Step 2: Center the weights with respect to the reference const centeredWeights = Array.from(currentContributions.values()).map((w) => @@ -120,10 +115,18 @@ export class PercentileClippingAggregator extends MultiRoundAggregator { clippedAvg.dispose(); // Step 7: Store result for next round - this.prevAggregate = result; + // Keep our own copy, the returned aggregate is owned (and disposed) by the caller + this.prevAggregate?.dispose(); + this.prevAggregate = result.clone(); return result; } + override dispose(): void { + this.prevAggregate?.dispose(); + this.prevAggregate = null; + super.dispose(); + } + private computePercentile(array: number[], percentile: number): number { // Linear interpolation for percentile calculation const clean = array.filter(Number.isFinite); From 56912792ace74677ec621c13953654424687341f Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Fri, 25 Sep 2026 16:07:41 +0200 Subject: [PATCH 04/23] fix: secure_history aggregator tensor leaks --- discojs/src/aggregator/secure_history.spec.ts | 33 +++++++++++++++++++ discojs/src/aggregator/secure_history.ts | 20 ++++++++--- 2 files changed, 48 insertions(+), 5 deletions(-) diff --git a/discojs/src/aggregator/secure_history.spec.ts b/discojs/src/aggregator/secure_history.spec.ts index 357493681..326f0e720 100644 --- a/discojs/src/aggregator/secure_history.spec.ts +++ b/discojs/src/aggregator/secure_history.spec.ts @@ -159,4 +159,37 @@ describe("Secure history aggregator", function () { expect(secureHistory).to.be.closeTo(secure, 0.001), ); }); + + it("aggregation leaves no dangling tensors", async () => { + const baseline = tf.memory().numTensors; + + const aggregator = new SecureHistoryAggregator(100, 0.8); + const ids = ["a", "b"]; + aggregator.setNodes(Set(ids)); + + // three aggregation rounds, the last two apply momentum smoothing + for (let round = 0; round < 3; round++) { + // communication round 0 sums the shares, round 1 averages the partial sums + for ( + let communicationRound = 0; + communicationRound < 2; + communicationRound++ + ) { + const contributions = ids.map((_, i) => + WeightsContainer.of([i, round]), + ); + const p = aggregator.getPromiseForAggregation(); + ids.forEach((id, i) => + aggregator.add(id, contributions[i], round, communicationRound), + ); + // the caller owns the aggregate, disposing it must not break the next round + (await p).dispose(); + contributions.forEach((c) => c.dispose()); + } + } + + // the aggregator only keeps the previous aggregate + aggregator.dispose(); + expect(tf.memory().numTensors).to.equal(baseline); + }); }); diff --git a/discojs/src/aggregator/secure_history.ts b/discojs/src/aggregator/secure_history.ts index 6b14fb9e7..35a0a0e3c 100644 --- a/discojs/src/aggregator/secure_history.ts +++ b/discojs/src/aggregator/secure_history.ts @@ -1,3 +1,5 @@ +import * as tf from "@tensorflow/tfjs"; + import type { WeightsContainer } from "#weights/index"; import { avg } from "#weights/index"; @@ -47,19 +49,27 @@ export class SecureHistoryAggregator extends SecureAggregator { const contribAvg = avg(currentContributions.values()); if (this.prevAggregate === null) { - this.prevAggregate = contribAvg; + // Keep our own copy, the returned aggregate is owned (and disposed) by the caller + this.prevAggregate = contribAvg.clone(); return contribAvg; } const updatedMomentum = this.prevAggregate.mapWith( contribAvg, - (prevT, currT) => prevT.mul(this.beta).add(currT.mul(1 - this.beta)), + (prevT, currT) => + tf.tidy(() => prevT.mul(this.beta).add(currT.mul(1 - this.beta))), ); + contribAvg.dispose(); - // Dispose old tensors to avoid memory leaks - this.prevAggregate.weights.forEach((t) => t.dispose()); - this.prevAggregate = updatedMomentum; + this.prevAggregate.dispose(); + this.prevAggregate = updatedMomentum.clone(); return updatedMomentum; } + + override dispose(): void { + this.prevAggregate?.dispose(); + this.prevAggregate = null; + super.dispose(); + } } From 8b60182cd120ee45dd41d1bc662b10c98796edf1 Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Fri, 25 Sep 2026 18:13:23 +0200 Subject: [PATCH 05/23] fix: aggregators makePayloads return copies --- discojs/src/aggregator.spec.ts | 19 +++- discojs/src/aggregator/aggregator.ts | 1 + discojs/src/aggregator/byzantine.ts | 2 +- discojs/src/aggregator/index.ts | 3 +- discojs/src/aggregator/mean.ts | 2 +- discojs/src/aggregator/percentile_clipping.ts | 4 +- discojs/src/aggregator/secure.spec.ts | 31 +++++++ discojs/src/aggregator/secure.ts | 7 +- .../decentralized/decentralized_client.ts | 92 ++++++++++++------- .../src/client/federated/federated_client.ts | 20 ++-- 10 files changed, 128 insertions(+), 53 deletions(-) diff --git a/discojs/src/aggregator.spec.ts b/discojs/src/aggregator.spec.ts index efa3935a7..1879833e9 100644 --- a/discojs/src/aggregator.spec.ts +++ b/discojs/src/aggregator.spec.ts @@ -4,13 +4,20 @@ import { type Aggregator, MeanAggregator, SecureAggregator, + ByzantineRobustAggregator, + PercentileClippingAggregator, } from "#aggregator/index"; import type { NodeID } from "#client/index"; import { WeightsContainer } from "#weights/index"; const AGGREGATORS: Set<[name: string, new () => Aggregator]> = Set.of< new () => Aggregator ->(MeanAggregator, SecureAggregator).map((Aggregator) => [ +>( + MeanAggregator, + SecureAggregator, + ByzantineRobustAggregator, + PercentileClippingAggregator, +).map((Aggregator) => [ // MeanAggregator waits for 100% of the node's contributions by default Aggregator.name, Aggregator, @@ -122,11 +129,15 @@ export async function communicate( if (contribution === undefined) throw new Error(`no contribution for ${id}`); - for (const [to, payload] of agg.makePayloads(contribution)) - network.get(to)?.add(id, payload.clone(), aggregationRound, r); + for (const [to, payload] of agg.makePayloads(contribution)) { + network.get(to)?.add(id, payload, aggregationRound, r); + payload.dispose(); + } } - contributions = Map(await Promise.all(nextContributions)); + const aggregated = Map(await Promise.all(nextContributions)); + if (r > 0) contributions.forEach((c) => c.dispose()); + contributions = aggregated; } return contributions; diff --git a/discojs/src/aggregator/aggregator.ts b/discojs/src/aggregator/aggregator.ts index ed2ddb49a..eac6b2645 100644 --- a/discojs/src/aggregator/aggregator.ts +++ b/discojs/src/aggregator/aggregator.ts @@ -264,6 +264,7 @@ export abstract class Aggregator extends EventEmitter<{ /** * Constructs the payloads sent to other nodes as contribution. + * The payloads are owned by the caller, who has to dispose them. * @param base Object from which the payload is computed */ abstract makePayloads(base: WeightsContainer): Map; diff --git a/discojs/src/aggregator/byzantine.ts b/discojs/src/aggregator/byzantine.ts index 2b90c974e..c406c9bd9 100644 --- a/discojs/src/aggregator/byzantine.ts +++ b/discojs/src/aggregator/byzantine.ts @@ -180,7 +180,7 @@ export class ByzantineRobustAggregator extends MultiRoundAggregator { weights: WeightsContainer, ): Map { // Communicate our local weights to every other node, be it a peer or a server - return this.nodes.toMap().map(() => weights); + return this.nodes.toMap().map(() => weights.clone()); } } diff --git a/discojs/src/aggregator/index.ts b/discojs/src/aggregator/index.ts index 75be361d5..7a3e9f2a9 100644 --- a/discojs/src/aggregator/index.ts +++ b/discojs/src/aggregator/index.ts @@ -1,5 +1,6 @@ export { Aggregator } from "./aggregator.js"; export { MeanAggregator } from "./mean.js"; export { SecureAggregator } from "./secure.js"; - +export { ByzantineRobustAggregator } from "./byzantine.js"; +export { PercentileClippingAggregator } from "./percentile_clipping.js"; export { getAggregator } from "./get.js"; diff --git a/discojs/src/aggregator/mean.ts b/discojs/src/aggregator/mean.ts index 3d5854436..efa861371 100644 --- a/discojs/src/aggregator/mean.ts +++ b/discojs/src/aggregator/mean.ts @@ -57,6 +57,6 @@ export class MeanAggregator extends MultiRoundAggregator { weights: WeightsContainer, ): Map { // Communicate our local weights to every other node, be it a peer or a server - return this.nodes.toMap().map(() => weights); + return this.nodes.toMap().map(() => weights.clone()); } } diff --git a/discojs/src/aggregator/percentile_clipping.ts b/discojs/src/aggregator/percentile_clipping.ts index cfa26e3a1..bcae340ec 100644 --- a/discojs/src/aggregator/percentile_clipping.ts +++ b/discojs/src/aggregator/percentile_clipping.ts @@ -3,7 +3,7 @@ import * as tf from "@tensorflow/tfjs"; import { AggregationStep } from "#aggregator/aggregator"; import type { ThresholdType } from "#aggregator/multiround"; import { MultiRoundAggregator } from "#aggregator/multiround"; -import type { NodeID } from "#client/index"; +import type { NodeID } from "#client/types"; import type { WeightsContainer } from "#weights/index"; import { avg } from "#weights/index"; @@ -147,7 +147,7 @@ export class PercentileClippingAggregator extends MultiRoundAggregator { override makePayloads( weights: WeightsContainer, ): Map { - return this.nodes.toMap().map(() => weights); + return this.nodes.toMap().map(() => weights.clone()); } } diff --git a/discojs/src/aggregator/secure.spec.ts b/discojs/src/aggregator/secure.spec.ts index d822bcb30..9fe9c88b7 100644 --- a/discojs/src/aggregator/secure.spec.ts +++ b/discojs/src/aggregator/secure.spec.ts @@ -1,5 +1,6 @@ import { List, Map, Range, Set } from "immutable"; import { assert, describe, expect, it } from "vitest"; +import * as tf from "@tensorflow/tfjs"; import { communicate, setupNetwork, wsIntoArrays } from "#root/aggregator.spec"; import { sum, avg, WeightsContainer } from "#weights/index"; @@ -87,4 +88,34 @@ describe("secure aggregator", () => { )) expect(secure).to.be.closeTo(mean, 0.001); }); + + it("generating shares leaves only the shares", () => { + const secret = WeightsContainer.of([1, 2, 3], [4]); + const aggregator = new SecureAggregator(); + aggregator.setNodes(Set.of("a", "b", "c")); + + const baseline = tf.memory().numTensors; + const shares = aggregator.generateAllShares(secret); + shares.forEach((share) => share.dispose()); + expect(tf.memory().numTensors).to.equal(baseline); + secret.dispose(); + }); + + it("aggregation leaves no dangling tensors", async () => { + const baseline = tf.memory().numTensors; + + const network = setupNetwork(SecureAggregator); + const contributions = network.map((_, id) => + WeightsContainer.of([id.length], [1, 2]), + ); + const results = await communicate( + network.map((agg, id) => [agg, contributions.get(id)!]), + 0, + ); + + results.forEach((result) => result.dispose()); + contributions.forEach((contribution) => contribution.dispose()); + network.forEach((agg) => agg.dispose()); + expect(tf.memory().numTensors).to.equal(baseline); + }); }); diff --git a/discojs/src/aggregator/secure.ts b/discojs/src/aggregator/secure.ts index 713691647..a7e8dcbec 100644 --- a/discojs/src/aggregator/secure.ts +++ b/discojs/src/aggregator/secure.ts @@ -95,7 +95,7 @@ export class SecureAggregator extends Aggregator { } // Send our partial sum to every other nodes case 1: - return this.nodes.toMap().map(() => weights); + return this.nodes.toMap().map(() => weights.clone()); default: throw new Error("communication round is out of bounds"); } @@ -112,7 +112,10 @@ export class SecureAggregator extends Aggregator { .toList(); // The last share completes the sum - return shares.push(secret.sub(sum(shares))); + const sharesSum = sum(shares); + const lastShare = secret.sub(sharesSum); + sharesSum.dispose(); + return shares.push(lastShare); } /** diff --git a/discojs/src/client/decentralized/decentralized_client.ts b/discojs/src/client/decentralized/decentralized_client.ts index fb9499290..4f8da341e 100644 --- a/discojs/src/client/decentralized/decentralized_client.ts +++ b/discojs/src/client/decentralized/decentralized_client.ts @@ -476,60 +476,77 @@ export class DecentralizedClient extends Client<"decentralized"> { // A communication round's payload is the aggregation result of the previous communication round. The first // communication round simply sends our training result, i.e. model weights updates. This scheme allows for // the aggregator to define any complex multi-round aggregation mechanism. + // `weights` is owned by the caller, every later `result` is ours to dispose let result = weights; + const disposeResult = () => { + if (result !== weights) result.dispose(); + }; for ( let communicationRound = 0; communicationRound < this.aggregator.communicationRounds; communicationRound++ ) { const connections = this.#connections; - if (connections === undefined) + if (connections === undefined) { + disposeResult(); throw new Error("peer's connections is undefined"); + } // Generate our payloads for this communication round and send them to all ready connected peers const payloads = this.aggregator.makePayloads(result); - await Promise.all( + const sent = Promise.all( payloads .entrySeq() .map(async ([id, payload]) => { - if (id === this.ownId) { - // add our own contribution to the aggregator, which takes a copy - this.aggregator.add( - this.ownId, - payload, + try { + if (id === this.ownId) { + // add our own contribution to the aggregator, which takes a copy + this.aggregator.add( + this.ownId, + payload, + this.aggregator.round, + communicationRound, + ); + return; + } + + const peer = connections.get(id); + if (peer === undefined) return; + + const encoded = await weightsEncode(payload); + + const msg: messages.PeerMessage = { + type: MType.Payload, + peer: id, + aggregationRound: this.aggregator.round, + communicationRound, + payload: encoded, + }; + + peer.send(msg); + + debug( + `[${shortenId(this.ownId)}] send weight update to peer ${shortenId(msg.peer)}` + + ` for round (%d, %d)`, this.aggregator.round, communicationRound, ); - return; + } finally { + payload.dispose(); } - - const peer = connections.get(id); - if (peer === undefined) return; - - const encoded = await weightsEncode(payload); - - const msg: messages.PeerMessage = { - type: MType.Payload, - peer: id, - aggregationRound: this.aggregator.round, - communicationRound, - payload: encoded, - }; - - peer.send(msg); - - debug( - `[${shortenId(this.ownId)}] send weight update to peer ${shortenId(msg.peer)}` + - ` for round (%d, %d)`, - this.aggregator.round, - communicationRound, - ); }) .toArray(), ); + try { + await sent; + } catch (e) { + disposeResult(); + throw e; + } // Wait for aggregation before proceeding to the next communication round. // The current result will be used as payload for the eventual next communication round. + let aggregated: WeightsContainer; try { - result = await Promise.race([ + aggregated = await Promise.race([ this.aggregationResult, timeout( undefined, @@ -537,7 +554,10 @@ export class DecentralizedClient extends Client<"decentralized"> { ), ]); } catch (e) { - if (this.isDisconnected) return weights.clone(); + if (this.isDisconnected) { + disposeResult(); + return weights.clone(); + } debug( `[${shortenId(this.ownId)}] while waiting for aggregation: %o`, @@ -545,6 +565,9 @@ export class DecentralizedClient extends Client<"decentralized"> { ); break; } + // the previous result has been sent, it is replaced by the new aggregation + disposeResult(); + result = aggregated; // There is at least one communication round remaining if (communicationRound < this.aggregator.communicationRounds - 1) { @@ -552,7 +575,10 @@ export class DecentralizedClient extends Client<"decentralized"> { this.aggregationResult = this.aggregator.getPromiseForAggregation(); } } - return await this.aggregationResult; + const aggregationResult = await this.aggregationResult; + // on the normal path, the last result is the final aggregation itself + if (result !== aggregationResult) disposeResult(); + return aggregationResult; } /** diff --git a/discojs/src/client/federated/federated_client.ts b/discojs/src/client/federated/federated_client.ts index dca53bac6..dfb00314d 100644 --- a/discojs/src/client/federated/federated_client.ts +++ b/discojs/src/client/federated/federated_client.ts @@ -148,16 +148,16 @@ export class FederatedClient extends Client<"federated"> { this.saveAndEmit("updating model"); // Send our local contribution to the server // and receive the server global update for this round as an answer to our contribution - const payloadToServer = this.aggregator - .makePayloads(weights) - .get(SERVER_NODE_ID); - if (payloadToServer === undefined) - throw new Error("aggregator didn't make a payload for the server"); - + // the payloads are ours to dispose + const payloads = this.aggregator.makePayloads(weights); const round = this.aggregator.round; - // block-scope the encoded payload so the potentially large buffer can be GC'd - // while we await the server's response below - { + try { + const payloadToServer = payloads.get(SERVER_NODE_ID); + if (payloadToServer === undefined) + throw new Error("aggregator didn't make a payload for the server"); + + // block-scope the encoded payload so the potentially large buffer can be GC'd + // while we await the server's response below const payload = await weightsEncode(payloadToServer); debug( "[%s] encoded payload for round %d byteLength=%d", @@ -175,6 +175,8 @@ export class FederatedClient extends Client<"federated"> { // Need to await the resulting global model right after sending our local contribution // to make sure we don't miss it this.server.send(msg); + } finally { + payloads.forEach((p) => p.dispose()); } debug( `[${shortenId(this.ownId)}] sent its local update to the server for round ${round}`, From 7652de12bde25ff710e690d9bc6970e8265d0cd9 Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 15:23:57 +0200 Subject: [PATCH 06/23] fix: dispose decoded GPT-2 weights --- discojs/src/serialization/model.spec.ts | 11 +++++++++++ discojs/src/serialization/model.ts | 7 ++++++- 2 files changed, 17 insertions(+), 1 deletion(-) diff --git a/discojs/src/serialization/model.spec.ts b/discojs/src/serialization/model.spec.ts index fba531fbe..4eb6f1d01 100644 --- a/discojs/src/serialization/model.spec.ts +++ b/discojs/src/serialization/model.spec.ts @@ -158,4 +158,15 @@ describe("serialization", () => { ); assert.deepEqual(model.config, decoded.config); }); + + it("decoding a gpt-tfjs model leaves no dangling tensors", async () => { + const model = new GPT({ modelType: "gpt-nano", contextLength: 8 }); + const encoded = await encode(model); + model.dispose(); + + const baseline = tf.memory().numTensors; + const decoded = await decode(encoded); + decoded.dispose(); + expect(tf.memory().numTensors).to.equal(baseline); + }); }); diff --git a/discojs/src/serialization/model.ts b/discojs/src/serialization/model.ts index 4912b4897..fe85dd465 100644 --- a/discojs/src/serialization/model.ts +++ b/discojs/src/serialization/model.ts @@ -144,7 +144,12 @@ export async function decode(encoded: Encoded): Promise> { "invalid encoding, gpt-tfjs model weights should be an encoding of its weights", ); const weights = w_decode(rawModel); - return GPT.deserialize({ weights, config }); + // the model's variables take their own reference, free the decoded ones + try { + return GPT.deserialize({ weights, config }); + } finally { + weights.dispose(); + } } default: throw new Error("invalid encoding, model type unrecognized"); From c071559ada1462a6c9d92288b2382b407ca50856 Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 15:24:18 +0200 Subject: [PATCH 07/23] fix: dispose attention mask and token embeddings --- .../models/implementations/gpt/gpt.spec.ts | 11 ++++++ .../src/models/implementations/gpt/layers.ts | 37 +++++++++++++++---- 2 files changed, 41 insertions(+), 7 deletions(-) diff --git a/discojs/src/models/implementations/gpt/gpt.spec.ts b/discojs/src/models/implementations/gpt/gpt.spec.ts index 05f3fe957..115b49969 100644 --- a/discojs/src/models/implementations/gpt/gpt.spec.ts +++ b/discojs/src/models/implementations/gpt/gpt.spec.ts @@ -1,4 +1,5 @@ import { List } from "immutable"; +import * as tf from "@tensorflow/tfjs"; import { describe, expect, it } from "vitest"; import type { DataFormat } from "#types/index"; @@ -45,4 +46,14 @@ describe("gpt-tfjs", () => { expect(input + output).equal(data); // Assert that the model completes 'Lorem ipsum dolor' with 'sit' }); + + it("frees all its tensors on dispose", () => { + const baseline = tf.memory().numTensors; + + const model = new GPT({ modelType: "gpt-nano", contextLength: 8 }); + // the shared token embedding and the attention masks must be freed too + model.dispose(); + + expect(tf.memory().numTensors).to.equal(baseline); + }); }); diff --git a/discojs/src/models/implementations/gpt/layers.ts b/discojs/src/models/implementations/gpt/layers.ts index 54c840a53..92670939a 100644 --- a/discojs/src/models/implementations/gpt/layers.ts +++ b/discojs/src/models/implementations/gpt/layers.ts @@ -98,7 +98,7 @@ export class CausalSelfAttention extends tf.layers.Layer { private readonly attnDrop: number; private readonly residDrop: number; private readonly seed: number; - private readonly mask: tf.Tensor2D; + private mask?: tf.Tensor2D; cAttnWeight?: tf.LayerVariable; cAttnBias?: tf.LayerVariable; cProjWeight?: tf.LayerVariable; @@ -117,18 +117,20 @@ export class CausalSelfAttention extends tf.layers.Layer { this.attnDrop = config.attnDrop; this.residDrop = config.residDrop; this.seed = config.seed; + } + override build(): void { // mask is a lower triangular matrix filled with 1 // calling bandPart zero out the upper triangular part of the all-ones matrix // from the doc: tf.linalg.band_part(input, -1, 0) ==> Lower triangular part - this.mask = tf.linalg.bandPart( - tf.ones([config.contextLength, config.contextLength]), - -1, - 0, + // it is owned by the layer, keep it from any enclosing tf.tidy + const { contextLength } = this.config; + this.mask = tf.keep( + tf.tidy(() => + tf.linalg.bandPart(tf.ones([contextLength, contextLength]), -1, 0), + ), ); - } - override build(): void { // key, query, value projections for all heads, but in a batch this.cAttnWeight = this.addWeight( "c_attn/kernel", @@ -171,6 +173,13 @@ export class CausalSelfAttention extends tf.layers.Layer { ); } + // called by tfjs once the layer isn't referenced anymore + protected override disposeWeights(): number { + this.mask?.dispose(); + this.mask = undefined; + return super.disposeWeights(); + } + override computeOutputShape( inputShape: tf.Shape | tf.Shape[], ): tf.Shape | tf.Shape[] { @@ -264,6 +273,8 @@ export class CausalSelfAttention extends tf.layers.Layer { } public applyCausalMask(att: tf.Tensor, T: number): tf.Tensor { + if (this.mask === undefined) + throw new Error("Model not built, mask is undefined"); // mask is lower triangular matrix filled with 1 const mask = this.mask.slice([0, 0], [T, T]); // 1 - mask => upper triangular matrix filled with 1 @@ -498,6 +509,18 @@ export class LMEmbedding extends tf.layers.Layer { ); } + /** + * tfjs counts one reference per application of the layer, and this layer is + * applied twice (token embedding and language modeling head). As the model + * disposes each of its layers only once, the shared embedding would never be + * freed. It is never used outside of its model, so release every reference. + */ + override dispose(): ReturnType { + let result = super.dispose(); + while (result.refCountAfterDispose > 0) result = super.dispose(); + return result; + } + override computeOutputShape( inputShape: tf.Shape | tf.Shape[], ): tf.Shape | tf.Shape[] { From 8ec0248ca7d39f275d2d4c035368007b4d0847de Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 15:39:47 +0200 Subject: [PATCH 08/23] fix: dispose model after serialization --- server/src/model_set.ts | 7 ++++++- server/tests/model_set.spec.ts | 27 +++++++++++++++++++++++++++ 2 files changed, 33 insertions(+), 1 deletion(-) create mode 100644 server/tests/model_set.spec.ts diff --git a/server/src/model_set.ts b/server/src/model_set.ts index ad2fc9aca..7e944b377 100644 --- a/server/src/model_set.ts +++ b/server/src/model_set.ts @@ -82,8 +82,13 @@ export class ModelSet extends EventEmitter<{ let encodedModel: EncodedModel; if (!Array.isArray(newModel)) { + // the model is only built to be encoded, the server keeps the encoding const model = await newModel.getModel(); - encodedModel = await modelEncode(model); + try { + encodedModel = await modelEncode(model); + } finally { + model.dispose(); + } } else { const model = newModel[1]; if (isEncoded(model)) { diff --git a/server/tests/model_set.spec.ts b/server/tests/model_set.spec.ts new file mode 100644 index 000000000..d6dace986 --- /dev/null +++ b/server/tests/model_set.spec.ts @@ -0,0 +1,27 @@ +import * as tf from "@tensorflow/tfjs-node"; +import { describe, expect, it } from "vitest"; + +import type { ModelCard } from "@epfml/discojs"; +import { GPT } from "@epfml/discojs"; + +import { ModelSet } from "../src/model_set.js"; + +describe("model set", () => { + it("doesn't keep the models built from cards", async () => { + // Tensorflow.js itself leaks the optimizer's iteration count + // inside tfjs-layers' LayersModel.save + // so this test uses GPT-nano + const card: ModelCard<"text"> = { + card: { id: "test-gpt", name: "test GPT", dataType: "text" }, + getModel: () => + Promise.resolve(new GPT({ modelType: "gpt-nano", contextLength: 8 })), + }; + const baseline = tf.memory().numTensors; + + const modelSet = new ModelSet(); + await modelSet.addModel(card); + + expect(modelSet.models.has(card.card.id)).to.be.true; + expect(tf.memory().numTensors).to.equal(baseline); + }); +}); From fa65be480f6b0ff49e7d1dc68f8cd45db35e3808 Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 15:48:21 +0200 Subject: [PATCH 09/23] fix: modelSync leak --- discojs/src/client/client.ts | 5 +++++ .../client/decentralized/decentralized_client.ts | 6 ++++-- discojs/src/training/disco.ts | 5 +++++ server/tests/e2e/decentralized.spec.ts | 16 +++++++++++++++- server/tests/e2e/helpers.ts | 4 +++- 5 files changed, 32 insertions(+), 4 deletions(-) diff --git a/discojs/src/client/client.ts b/discojs/src/client/client.ts index ae825750b..eb92d16a6 100644 --- a/discojs/src/client/client.ts +++ b/discojs/src/client/client.ts @@ -24,6 +24,11 @@ const debug = createDebug("discojs:client"); export abstract class Client extends EventEmitter<{ status: RoundStatus; participants: number; + /** + * The latest global model received from another node. + * The weights are lent to the listeners for the duration of the call, + * clone them to keep them longer. + */ modelSynced: WeightsContainer; }> { // Own ID provided by the network's server. diff --git a/discojs/src/client/decentralized/decentralized_client.ts b/discojs/src/client/decentralized/decentralized_client.ts index 4f8da341e..fbe910d64 100644 --- a/discojs/src/client/decentralized/decentralized_client.ts +++ b/discojs/src/client/decentralized/decentralized_client.ts @@ -233,10 +233,12 @@ export class DecentralizedClient extends Client<"decentralized"> { const latestModel = await this.receiveModel(providerConn); + // keep the decoded model, listeners only borrow it + // emitting a clone would be very hard to dispose safely this.#latestModel?.dispose(); - this.#latestModel = this.cloneWeights(latestModel); + this.#latestModel = latestModel; - this.emit("modelSynced", this.cloneWeights(latestModel)); + this.emit("modelSynced", latestModel); this.#modelSyncNeeded = false; } diff --git a/discojs/src/training/disco.ts b/discojs/src/training/disco.ts index 278396f93..c7dddc987 100644 --- a/discojs/src/training/disco.ts +++ b/discojs/src/training/disco.ts @@ -69,6 +69,11 @@ function buildSummaryLog( export class Disco extends EventEmitter<{ status: RoundStatus; participants: number; + /** + * The model was synced to the latest global model. + * The weights are lent to the listeners for the duration of the call, + * clone them to keep them longer. + */ modelSynced: WeightsContainer | undefined; }> { public readonly trainer: Trainer; diff --git a/server/tests/e2e/decentralized.spec.ts b/server/tests/e2e/decentralized.spec.ts index a6293902d..31fd4fa86 100644 --- a/server/tests/e2e/decentralized.spec.ts +++ b/server/tests/e2e/decentralized.spec.ts @@ -689,6 +689,9 @@ describe("end-to-end decentralized", { timeout: 50_000 }, () => { const url = await startServer(defaultModels.LUSClassifier, taskProvider); const dataset = await datasets.loadLusCOVID(); + // the server and dataset are initialized, but not the clients + const memoryBeforeClients = tensorMemorySnapshot(); + const discoUser1 = new Disco(task, url, { preprocessOnce: true }); const discoUser2 = new Disco(task, url, { preprocessOnce: true }); const discoUser3 = new Disco(task, url, { preprocessOnce: true }); @@ -721,7 +724,8 @@ describe("end-to-end decentralized", { timeout: 50_000 }, () => { const waitForModelSynced = Promise.race([ new Promise((resolve) => { discoUser3.on("modelSynced", (weights) => { - if (weights !== undefined) resolve(weights); + // the event only lends the weights + if (weights !== undefined) resolve(weights.clone()); }); }), new Promise((_, reject) => @@ -769,12 +773,22 @@ describe("end-to-end decentralized", { timeout: 50_000 }, () => { expect(user3Round.done).to.be.false; boundaryAfterLastRound.forEach((model) => model.dispose()); + syncedWeights.dispose(); } finally { + // release the recorded models, they are not part of what we measure + modelsUser1.dispose(); + modelsUser2.dispose(); // Close clients if not already done await discoUser1.close().catch(() => {}); await discoUser2.close().catch(() => {}); await discoUser3.close().catch(() => {}); } + + // the model received by the newcomer must be released with the clients + await new Promise((resolve) => setImmediate(resolve)); + expect(tensorMemorySnapshot().numTensors).to.be.at.most( + memoryBeforeClients.numTensors, + ); }, ); diff --git a/server/tests/e2e/helpers.ts b/server/tests/e2e/helpers.ts index 4d79981a1..828f7f64d 100644 --- a/server/tests/e2e/helpers.ts +++ b/server/tests/e2e/helpers.ts @@ -158,7 +158,8 @@ export class Participant { ); this.#syncedModel = new Promise((resolve) => this.disco.on("modelSynced", (weights) => { - if (weights !== undefined) resolve(weights); + // the event only lends the weights + if (weights !== undefined) resolve(weights.clone()); }), ); @@ -232,6 +233,7 @@ export class Participant { await this.disco.close(); } finally { this.modelsAtRoundBoundary.dispose(); + void this.#syncedModel.then((weights) => weights.dispose()); } } } From a1064cc12aca497535b52011c0f3b27881ce218a Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 15:49:52 +0200 Subject: [PATCH 10/23] fix: dispose lossTensor --- discojs/src/models/implementations/hellaswag.ts | 1 + 1 file changed, 1 insertion(+) diff --git a/discojs/src/models/implementations/hellaswag.ts b/discojs/src/models/implementations/hellaswag.ts index aeedf0545..7f14323b7 100644 --- a/discojs/src/models/implementations/hellaswag.ts +++ b/discojs/src/models/implementations/hellaswag.ts @@ -76,6 +76,7 @@ async function computeLogLikelihood( return loss; }); const lossNumber = await lossTensor.array(); + lossTensor.dispose(); if (typeof lossNumber !== "number") { throw new Error("got multiple loss"); } From fc24e86f9bec50fe8a06e1f072ff6ee0d81369f7 Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 16:17:22 +0200 Subject: [PATCH 11/23] fix: avoid calling already disposed model --- docs/examples/wikitext.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/examples/wikitext.ts b/docs/examples/wikitext.ts index 742417da6..20da31417 100644 --- a/docs/examples/wikitext.ts +++ b/docs/examples/wikitext.ts @@ -38,7 +38,7 @@ async function main(): Promise { await disco.trainFully(dataset); // Get the model and save the trained model - model = disco.trainer.model as GPT; + model = disco.trainer.releaseModel() as GPT; await saveModelToDisk(model, modelFolder, modelFileName); await disco.close(); } else { From 1f85d0c731a6698dffd96a3d125922825f500f37 Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 16:17:39 +0200 Subject: [PATCH 12/23] doc: use trainer.releaseModel --- discojs/src/training/disco.ts | 3 +++ 1 file changed, 3 insertions(+) diff --git a/discojs/src/training/disco.ts b/discojs/src/training/disco.ts index c7dddc987..6765fb716 100644 --- a/discojs/src/training/disco.ts +++ b/discojs/src/training/disco.ts @@ -310,6 +310,9 @@ export class Disco extends EventEmitter<{ /** * Completely stops the ongoing training instance. + * Disposes all tensors including the model, + * call disco.trainer.releaseModel() if you need the model + * after closing disco */ async close(): Promise { // Dispose the model tensor and the aggregator's buffered tensors From 5cbb4e66d5fde8ea153c0c8fd360a82dea9cc934 Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 16:23:00 +0200 Subject: [PATCH 13/23] fix: stop inference when component unmounts --- .../src/components/testing/PredictSteps.vue | 5 +++- webapp/src/components/testing/TestSteps.vue | 30 +++++++++++-------- 2 files changed, 21 insertions(+), 14 deletions(-) diff --git a/webapp/src/components/testing/PredictSteps.vue b/webapp/src/components/testing/PredictSteps.vue index ee5798d72..1c6e488b6 100644 --- a/webapp/src/components/testing/PredictSteps.vue +++ b/webapp/src/components/testing/PredictSteps.vue @@ -100,7 +100,7 @@ import * as d3 from "d3"; import createDebug from "debug"; import { List } from "immutable"; -import { computed, ref, toRaw } from "vue"; +import { computed, onUnmounted, ref, toRaw } from "vue"; import type { DataFormat, @@ -263,6 +263,9 @@ async function startTabularInference( } } +// the model may be disposed once we're gone, don't keep inferring with it +onUnmounted(() => void stopInference()); + async function stopInference(): Promise { const g = generator.value; if (g === undefined) return; diff --git a/webapp/src/components/testing/TestSteps.vue b/webapp/src/components/testing/TestSteps.vue index 120ce8c4b..ff2c621ad 100644 --- a/webapp/src/components/testing/TestSteps.vue +++ b/webapp/src/components/testing/TestSteps.vue @@ -152,7 +152,7 @@ import * as d3 from "d3"; import createDebug from "debug"; import { List, Map } from "immutable"; -import { computed, ref, toRaw } from "vue"; +import { computed, onUnmounted, ref, toRaw } from "vue"; import type { DataType, Image, Model, Network, Task } from "@epfml/discojs"; import { Validator } from "@epfml/discojs"; @@ -346,9 +346,9 @@ async function startImageTest( const validator = new Validator(task, model); let results: Tested["image"] = List(); + // stopTest clears controller, keep our own signal + const { signal } = (controller.value = new AbortController()); try { - controller.value = new AbortController(); - for await (const [ { filename, image, label }, { predicted, truth }, @@ -376,10 +376,10 @@ async function startImageTest( tested.value = results as Tested[D]; - if (controller.value.signal.aborted) break; + if (signal.aborted) break; } } finally { - controller.value = undefined; + if (controller.value?.signal === signal) controller.value = undefined; } } @@ -401,9 +401,9 @@ async function startTabularTest( const validator = new Validator(task, model); let results: Tested["tabular"]["results"] = List(); + // stopTest clears controller, keep our own signal + const { signal } = (controller.value = new AbortController()); try { - controller.value = new AbortController(); - for await (const [row, { predicted, truth }] of dataset.zip( validator.test(dataset), )) { @@ -428,10 +428,10 @@ async function startTabularTest( tested.value = { labels, results } as Tested[D]; - if (controller.value.signal.aborted) break; + if (signal.aborted) break; } } finally { - controller.value = undefined; + if (controller.value?.signal === signal) controller.value = undefined; } } @@ -443,25 +443,29 @@ async function startTextTest( const validator = new Validator(task, model); let results: Tested["text"] = List(); + // stopTest clears controller, keep our own signal + const { signal } = (controller.value = new AbortController()); try { - controller.value = new AbortController(); - for await (const { predicted, truth } of validator.test(dataset)) { results = results.push({ output: { correct: predicted === truth } }); tested.value = results as Tested[D]; - if (controller.value.signal.aborted) break; + if (signal.aborted) break; // TODO processing can hog the browser when big enough // this allow other computations to run // will be fixed by using WebWorker await new Promise((resolve) => setTimeout(() => resolve(), 100)); + if (signal.aborted) break; } } finally { - controller.value = undefined; + if (controller.value?.signal === signal) controller.value = undefined; } } +// the model may be disposed once we're gone, don't keep testing it +onUnmounted(() => stopTest()); + function stopTest(): void { const c = controller.value; if (c === undefined) return; From e0352f5fd2b3a3cdddd3cf95a6ce00fc723a815d Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 16:23:13 +0200 Subject: [PATCH 14/23] fix: dispose model after creating new task --- webapp/src/components/task_creation_form/TaskCreationForm.vue | 3 +++ 1 file changed, 3 insertions(+) diff --git a/webapp/src/components/task_creation_form/TaskCreationForm.vue b/webapp/src/components/task_creation_form/TaskCreationForm.vue index efa4fde0c..aade49b33 100644 --- a/webapp/src/components/task_creation_form/TaskCreationForm.vue +++ b/webapp/src/components/task_creation_form/TaskCreationForm.vue @@ -1169,6 +1169,9 @@ async function onSubmit(form: unknown): Promise { toaster.error("This identifier is already taken"); else toaster.error("An error occured server-side"); return; + } finally { + // the server keeps its own copy + model.dispose(); } if (typeof tasks.value === "string") From c45239b3d1089284884c1252ce4480d33e2fb36b Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 16:23:31 +0200 Subject: [PATCH 15/23] fix: dispose llm when leaving ChatUI --- webapp/src/components/testing/ChatUI.vue | 62 +++++++++++++++++++++--- 1 file changed, 54 insertions(+), 8 deletions(-) diff --git a/webapp/src/components/testing/ChatUI.vue b/webapp/src/components/testing/ChatUI.vue index 09c1d0fa3..a75355da9 100644 --- a/webapp/src/components/testing/ChatUI.vue +++ b/webapp/src/components/testing/ChatUI.vue @@ -338,8 +338,10 @@ From c42c5530059e23152f3d6032ef2fedc27a97fb1f Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 16:24:02 +0200 Subject: [PATCH 16/23] fix: dispose unused models --- .../src/components/testing/ModelLibrary.vue | 44 ++++++++++++++++--- 1 file changed, 38 insertions(+), 6 deletions(-) diff --git a/webapp/src/components/testing/ModelLibrary.vue b/webapp/src/components/testing/ModelLibrary.vue index 47adb7451..cf3e9026b 100644 --- a/webapp/src/components/testing/ModelLibrary.vue +++ b/webapp/src/components/testing/ModelLibrary.vue @@ -135,13 +135,16 @@
+ @@ -155,7 +158,7 @@ import createDebug from "debug"; import { List } from "immutable"; import { storeToRefs } from "pinia"; -import { computed, ref, onActivated } from "vue"; +import { computed, nextTick, ref, onActivated, onUnmounted } from "vue"; import { RouterLink } from "vue-router"; import { VueSpinner } from "vue3-spinners"; @@ -188,6 +191,7 @@ const toaster = useToaster(); const router = useRouter(); type Selection = { + modelID: ModelID; mode: "predict" | "test"; task: Task; // same as in validation store but not undef @@ -242,11 +246,24 @@ function formatByteSize(size: number): string { }).format(size); } +// incremented on every selection and on unmount, so that a selection +// finishing late knows its model isn't wanted anymore +let selectionID = 0; + +onUnmounted(() => { + selectionID++; + selection.value?.model.dispose(); + selection.value = undefined; +}); + onActivated(() => { // handle test after training or from library // TODO encode model ID inside the URL instead of relying on store - if (validationStore.modelID !== undefined) - void selectModel(validationStore.modelID, "test"); + const { modelID } = validationStore; + // keep the current selection, a test or an inference may still be running + // on it since the page was left + if (modelID !== undefined && modelID !== selection.value?.modelID) + void selectModel(modelID, "test"); }); async function downloadModel(task: Task): Promise { @@ -265,8 +282,12 @@ async function downloadModel(task: Task): Promise { getAggregator(task), ); const model = await client.getLatestModel(); - - await models.add(task.id, model); + // the store keeps an encoded copy + try { + await models.add(task.id, model); + } finally { + model.dispose(); + } } catch (e) { debug("while downloading model: %o", e); toaster.error("Something went wrong, please try again later."); @@ -288,8 +309,15 @@ async function selectModel( modelID: ModelID, mode: "predict" | "test", ): Promise { + const id = ++selectionID; + // decoded for us, so it's ours to dispose const model = await models.get(modelID); if (model === undefined) throw new Error("model ID not present in store"); + // a newer selection started or the page was unmounted in the meantime + if (id !== selectionID) { + model.dispose(); + return; + } const taskID = models.infos.get(modelID)?.taskID; if (taskID === undefined) throw new Error("task ID for model ID not found"); @@ -298,10 +326,14 @@ async function selectModel( const task = tasks.value.get(taskID); if (task === undefined) throw new Error("task not found"); - selection.value = { mode, model, task }; + const previous = selection.value; + selection.value = { modelID, mode, model, task }; validationStore.mode = mode; validationStore.modelID = modelID; validationStore.step = 1; + // once re-rendered, the steps using the previous model are unmounted + await nextTick(); + previous?.model.dispose(); } function removeModel(modelID: ModelID): void { From 2db3237ef38b036a6da6a36b9a59fb49b759ae55 Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 16:46:04 +0200 Subject: [PATCH 17/23] fix: dispose on error --- cli/src/args.ts | 33 +++-- cli/src/evaluate_finetuned_gpt2.ts | 26 ++-- cli/src/measure_memorization_gpt2.ts | 24 ++-- discojs/src/models/implementations/gpt/gpt.ts | 39 +++--- .../src/models/implementations/gpt/model.ts | 91 +++++++------ discojs/src/models/tfjs.ts | 125 +++++++++--------- discojs/src/privacy.ts | 44 +++--- discojs/src/training/trainer.ts | 70 ++++++---- docs/examples/training.ts | 12 +- .../task_creation_form/TaskCreationForm.vue | 17 ++- 10 files changed, 266 insertions(+), 215 deletions(-) diff --git a/cli/src/args.ts b/cli/src/args.ts index 34e35100a..a0bc4d9e4 100644 --- a/cli/src/args.ts +++ b/cli/src/args.ts @@ -445,21 +445,26 @@ export const args: BenchmarkArguments = { async getModel() { const model = await provider.modelCard.getModel(); - if (unsafeArgs.learningRate !== undefined) { - if (!(model instanceof GPT)) - throw new Error( - "learningRate override is only supported for GPT models", + try { + if (unsafeArgs.learningRate !== undefined) { + if (!(model instanceof GPT)) + throw new Error( + "learningRate override is only supported for GPT models", + ); + if ( + !Number.isFinite(unsafeArgs.learningRate) || + unsafeArgs.learningRate <= 0 + ) + throw new Error("learningRate must be a positive finite number"); + + model.setLearningRate(unsafeArgs.learningRate); + console.log( + `Overriding GPT learning rate to ${unsafeArgs.learningRate}`, ); - if ( - !Number.isFinite(unsafeArgs.learningRate) || - unsafeArgs.learningRate <= 0 - ) - throw new Error("learningRate must be a positive finite number"); - - model.setLearningRate(unsafeArgs.learningRate); - console.log( - `Overriding GPT learning rate to ${unsafeArgs.learningRate}`, - ); + } + } catch (e) { + model.dispose(); + throw e; } return model; diff --git a/cli/src/evaluate_finetuned_gpt2.ts b/cli/src/evaluate_finetuned_gpt2.ts index b3409c657..c17d291dd 100644 --- a/cli/src/evaluate_finetuned_gpt2.ts +++ b/cli/src/evaluate_finetuned_gpt2.ts @@ -93,6 +93,7 @@ function predictTokenLogits( ): tf.Tensor3D { const logits = tfModel.predict(inputTensor); if (Array.isArray(logits)) { + tf.dispose(logits); throw new Error("Expected GPT model to return a single logits tensor"); } if (logits.rank !== 3) { @@ -244,12 +245,6 @@ async function scoreContinuations( ...Array(maxInputLength - truncatedInputTokens.length).fill(0), ]); - const inputTensor = tf.tensor2d( - paddedInputs, - [paddedInputs.length, maxInputLength], - "int32", - ); - const targetIndexes: number[][] = []; const targetTokenIds: number[] = []; const targetOwners: number[] = []; @@ -275,7 +270,6 @@ async function scoreContinuations( }); if (targetIndexes.length === 0) { - inputTensor.dispose(); return scoredInputs.map((scoredInput) => ({ score: Number.NEGATIVE_INFINITY, promptTokens: promptTokens.length, @@ -284,8 +278,13 @@ async function scoreContinuations( })); } - const logits = predictTokenLogits(tfModel, inputTensor); const targetLogProbs = tf.tidy(() => { + const inputTensor = tf.tensor2d( + paddedInputs, + [paddedInputs.length, maxInputLength], + "int32", + ); + const logits = predictTokenLogits(tfModel, inputTensor); const targetIndexTensor = tf.tensor2d( targetIndexes, [targetIndexes.length, 2], @@ -301,7 +300,12 @@ async function scoreContinuations( return tf.gatherND(logProbs, targetTokenIndexTensor); }); - const targetScores = (await targetLogProbs.array()) as number[]; + let targetScores: number[]; + try { + targetScores = (await targetLogProbs.array()) as number[]; + } finally { + targetLogProbs.dispose(); + } const scoreSums = Array(scoredInputs.length).fill(0) as number[]; const scoreCounts = Array(scoredInputs.length).fill(0) as number[]; @@ -321,10 +325,6 @@ async function scoreContinuations( usedInputTokens: scoredInput.truncatedInputTokens.length, })); - inputTensor.dispose(); - logits.dispose(); - targetLogProbs.dispose(); - return results; } diff --git a/cli/src/measure_memorization_gpt2.ts b/cli/src/measure_memorization_gpt2.ts index 0187a87ad..c3dc31415 100644 --- a/cli/src/measure_memorization_gpt2.ts +++ b/cli/src/measure_memorization_gpt2.ts @@ -187,17 +187,11 @@ async function sampleGenerateGPT2( for (let i = 0; i < maxNewTokens; i++) { const modelInput = generated.slice(-maxContextLength); - const input = tf.tensor2d([modelInput], [1, modelInput.length], "int32"); - - const logits = tf.tidy(() => { + const nextTokenTensor = tf.tidy(() => { + const input = tf.tensor2d([modelInput], [1, modelInput.length], "int32"); const output = tfModel.predict(input); - if (Array.isArray(output)) { - return output[0]; - } - return output; - }); + const logits = Array.isArray(output) ? output[0] : output; - const nextTokenTensor = tf.tidy(() => { const last = logits.slice([0, modelInput.length - 1, 0], [1, 1, -1]); const scaled = last.squeeze().div(temperature); const { values: topKLogits, indices: topKTokens } = tf.topk(scaled, topK); @@ -213,12 +207,12 @@ async function sampleGenerateGPT2( return topKTokens.gather(sampledIndex).squeeze(); }); - const nextTokenData = await nextTokenTensor.data(); - const nextToken = nextTokenData[0]; - - input.dispose(); - logits.dispose(); - nextTokenTensor.dispose(); + let nextToken: number; + try { + nextToken = (await nextTokenTensor.data())[0]; + } finally { + nextTokenTensor.dispose(); + } generated.push(nextToken); } diff --git a/discojs/src/models/implementations/gpt/gpt.ts b/discojs/src/models/implementations/gpt/gpt.ts index 280895038..47fcf9bef 100644 --- a/discojs/src/models/implementations/gpt/gpt.ts +++ b/discojs/src/models/implementations/gpt/gpt.ts @@ -140,17 +140,20 @@ export class GPT extends Model<"text"> { const tfBatch = this.#batchToTF(batch); let logs: tf.Logs | undefined; - await this.model.fitDataset(tf.data.array([tfBatch]), { - epochs: 1, - iterationOffset: iterationNumber - 1, - verbose: 0, // don't pollute - callbacks: { - onEpochEnd: (_, cur) => { - logs = cur; + try { + await this.model.fitDataset(tf.data.array([tfBatch]), { + epochs: 1, + iterationOffset: iterationNumber - 1, + verbose: 0, // don't pollute + callbacks: { + onEpochEnd: (_, cur) => { + logs = cur; + }, }, - }, - }); - tf.dispose(tfBatch); + }); + } finally { + tf.dispose(tfBatch); + } if (logs === undefined) throw new Error("batch didn't gave any logs"); const { loss, acc: accuracy } = logs; @@ -230,18 +233,16 @@ export class GPT extends Model<"text"> { // slice input tokens if longer than context length tokens = tokens.slice(-this.#contextLength); - const input = tf.tidy(() => - tf.tensor1d(tokens.toArray(), "int32").expandDims(0), - ); - const logits = tf.tidy(() => { + const input = tf + .tensor1d(tokens.toArray(), "int32") + .expandDims(0); const output = this.model.predict(input); if (Array.isArray(output)) throw new Error("The model outputs too multiple values"); if (output.rank !== 3) throw new Error("The model outputs wrong shape"); return output.squeeze([0]); }); - input.dispose(); const probs = tf.tidy(() => logits @@ -281,9 +282,11 @@ export class GPT extends Model<"text"> { }); probs.dispose(); - const ret = await next.array(); - next.dispose(); - return ret; + try { + return await next.array(); + } finally { + next.dispose(); + } } get config(): Required { diff --git a/discojs/src/models/implementations/gpt/model.ts b/discojs/src/models/implementations/gpt/model.ts index 8ff763200..d99cb0957 100644 --- a/discojs/src/models/implementations/gpt/model.ts +++ b/discojs/src/models/implementations/gpt/model.ts @@ -141,13 +141,8 @@ export class GPTModel extends tf.LayersModel { while (next.done !== true && iteration <= this.config.maxIter) { const reportedIteration = iterationOffset + iteration; let weightUpdateTime = performance.now(); - await callbacks.onEpochBegin?.(epoch); const { xs, ys } = next.value as { xs: tf.Tensor2D; ys: tf.Tensor3D }; - let preprocessingTime = performance.now(); - await Promise.all([xs.data(), ys.data()]); - preprocessingTime = performance.now() - preprocessingTime; - // TODO include as a tensor inside the model // const accTensor = tf.tidy(() => { // const logits = this.apply(xs) @@ -167,40 +162,58 @@ export class GPTModel extends tf.LayersModel { // tf.dispose([accTensor]) accuracyFraction = [Number.NaN, Number.NaN]; - const goldfishLoss = this.#goldfishLoss; - const goldfishMask = - goldfishLoss === undefined - ? undefined - : this.#buildGoldfishMask(xs, goldfishLoss); - - const lossTensor = tf.tidy(() => { - const { grads, value: lossTensor } = this.optimizer.computeGradients( - () => { - const logits = this.apply(xs); - if (Array.isArray(logits)) - throw new Error("model outputs too many tensor"); - if (logits instanceof tf.SymbolicTensor) - throw new Error("model outputs symbolic tensor"); - return goldfishMask === undefined || goldfishLoss === undefined - ? tf.losses.softmaxCrossEntropy(ys, logits) - : this.#goldfishLossTensor( - ys, - logits, - goldfishMask, - goldfishLoss, - ); - }, - ); - const gradsClipped = clipByGlobalNormObj(grads, 1); - this.optimizer.applyGradients(gradsClipped); - tf.dispose(Object.values(gradsClipped)); - return lossTensor; - }); - goldfishMask?.dispose(); - - const loss = await lossTensor.array(); - lossTensor.dispose(); - tf.dispose([xs, ys]); + let preprocessingTime: number; + let loss: number; + try { + await callbacks.onEpochBegin?.(epoch); + + preprocessingTime = performance.now(); + await Promise.all([xs.data(), ys.data()]); + preprocessingTime = performance.now() - preprocessingTime; + + const goldfishLoss = this.#goldfishLoss; + const goldfishMask = + goldfishLoss === undefined + ? undefined + : this.#buildGoldfishMask(xs, goldfishLoss); + + let lossTensor: tf.Scalar; + try { + lossTensor = tf.tidy(() => { + const { grads, value: lossTensor } = + this.optimizer.computeGradients(() => { + const logits = this.apply(xs); + if (Array.isArray(logits)) + throw new Error("model outputs too many tensor"); + if (logits instanceof tf.SymbolicTensor) + throw new Error("model outputs symbolic tensor"); + return goldfishMask === undefined || + goldfishLoss === undefined + ? tf.losses.softmaxCrossEntropy(ys, logits) + : this.#goldfishLossTensor( + ys, + logits, + goldfishMask, + goldfishLoss, + ); + }); + const gradsClipped = clipByGlobalNormObj(grads, 1); + this.optimizer.applyGradients(gradsClipped); + tf.dispose(Object.values(gradsClipped)); + return lossTensor; + }); + } finally { + goldfishMask?.dispose(); + } + + try { + loss = await lossTensor.array(); + } finally { + lossTensor.dispose(); + } + } finally { + tf.dispose([xs, ys]); + } averageLoss += loss; weightUpdateTime = performance.now() - weightUpdateTime; diff --git a/discojs/src/models/tfjs.ts b/discojs/src/models/tfjs.ts index eeffb523c..e6fa92ea4 100644 --- a/discojs/src/models/tfjs.ts +++ b/discojs/src/models/tfjs.ts @@ -104,17 +104,21 @@ export class TFJS extends Model { }.bind(this), ), ); - const metricToValue = Map( - List(this.model.metricsNames).zip( - Array.isArray(evaluation) - ? List(await Promise.all(evaluation.map((t) => t.data()))) - : List.of(await evaluation.data()), - ), - ).map((values) => { - if (values.length !== 1) throw new Error("more than one metric value"); - return values[0]; - }); - tf.dispose(evaluation); + let metricToValue: Map; + try { + metricToValue = Map( + List(this.model.metricsNames).zip( + Array.isArray(evaluation) + ? List(await Promise.all(evaluation.map((t) => t.data()))) + : List.of(await evaluation.data()), + ), + ).map((values) => { + if (values.length !== 1) throw new Error("more than one metric value"); + return values[0]; + }); + } finally { + tf.dispose(evaluation); + } const [accuracy, loss] = [ metricToValue.get("acc"), @@ -130,53 +134,52 @@ export class TFJS extends Model { batch: Batched, ): Promise> { async function cleanupPredicted(y: tf.Tensor1D): Promise { - if (y.shape[0] === 1) { - // Binary classification - const threshold = tf.scalar(0.5); - const binaryTensor = y.greaterEqual(threshold); - - const binaryArray = await binaryTensor.data(); - tf.dispose([y, binaryTensor, threshold]); - - return binaryArray[0]; - } - - // Multi-class classification - const indexTensor = y.argMax(); - - const indexArray = await indexTensor.data(); - tf.dispose([y, indexTensor]); - - return indexArray[0]; - + // Binary classification if single output, multi-class otherwise // Multi-label classification is not supported + const predicted = tf.tidy(() => + y.shape[0] === 1 ? y.greaterEqual(tf.scalar(0.5)) : y.argMax(), + ); + try { + return (await predicted.data())[0]; + } finally { + predicted.dispose(); + } } const xs = this.#batchWithoutLabelToTF(batch); + let prediction: tf.Tensor | tf.Tensor[]; + try { + prediction = this.model.predict(xs); + } finally { + tf.dispose(xs); + } - const prediction = this.model.predict(xs); - if (Array.isArray(prediction)) - throw new Error( - "prediction yield many Tensors but should have only returned one", - ); - tf.dispose(xs); - - if (prediction.rank !== 2) - throw new Error("unexpected batched prediction shape"); - - const ret = List( - await Promise.all( - tf.unstack(prediction).map((y) => - cleanupPredicted( - // cast as unstack reduce by one the rank - y as tf.Tensor1D, + try { + if (Array.isArray(prediction)) + throw new Error( + "prediction yield many Tensors but should have only returned one", + ); + if (prediction.rank !== 2) + throw new Error("unexpected batched prediction shape"); + + const ys = tf.unstack(prediction); + try { + return List( + await Promise.all( + ys.map((y) => + cleanupPredicted( + // cast as unstack reduce by one the rank + y as tf.Tensor1D, + ), + ), ), - ), - ), - ); - prediction.dispose(); - - return ret; + ); + } finally { + tf.dispose(ys); + } + } finally { + tf.dispose(prediction); + } } static async deserialize([ @@ -184,13 +187,17 @@ export class TFJS extends Model { artifacts, metadata, ]: Serialized): Promise> { - return new this( - datatype, - await tf.loadLayersModel({ - load: () => Promise.resolve(artifacts), - }), - metadata, - ); + const model = await tf.loadLayersModel({ + load: () => Promise.resolve(artifacts), + }); + try { + return new this(datatype, model, metadata); + } catch (e) { + // constructor validation failed, don't leak the loaded weights + model.dispose(); + model.optimizer?.dispose(); + throw e; + } } async serialize(): Promise> { diff --git a/discojs/src/privacy.ts b/discojs/src/privacy.ts index 3ddd76b0d..6fa8a1c89 100644 --- a/discojs/src/privacy.ts +++ b/discojs/src/privacy.ts @@ -7,10 +7,13 @@ import type { WeightNormHistory } from "#training/types"; /** Computes the Frobenius norm of the given weights. */ export async function frobeniusNorm(weights: tf.Tensor): Promise { const squaredTensor = tf.tidy(() => weights.square().sum()); - const squared = await squaredTensor.data(); - squaredTensor.dispose(); - if (squared.length !== 1) throw new Error("unexpected weights shape"); - return Math.sqrt(squared[0]); + try { + const squared = await squaredTensor.data(); + if (squared.length !== 1) throw new Error("unexpected weights shape"); + return Math.sqrt(squared[0]); + } finally { + squaredTensor.dispose(); + } } /** ALDP-FL implementation */ @@ -55,8 +58,12 @@ export async function addOptimalNoise( const clippedWeights = await clipNorm(weightUpdates, clippingRadius); try { - return clippedWeights.map((w, i) => - tf.tidy(() => w.add(tf.randomNormal(w.shape, 0, sigmas[i]))), + return new WeightsContainer( + tf.tidy(() => + clippedWeights.weights.map((w, i) => + w.add(tf.randomNormal(w.shape, 0, sigmas[i])), + ), + ), ); } finally { clippedWeights.dispose(); @@ -80,19 +87,18 @@ export async function clipNorm( `radius length mismatch: got ${radius.length}, expected ${layers.length}`, ); - /** Apply different clipping radius to each layer in the WeightsContainer */ - const clipped = await Promise.all( - layers.map(async (l, i) => { - const norm = await frobeniusNorm(l); - const r = radius[i]; + // Check the invalid radius value + if (radius.some((r) => !Number.isFinite(r) || r <= 0)) + throw new Error("Invalid radius value"); - // Check the invalid radius value - if (!Number.isFinite(r) || r <= 0) - throw new Error("Invalid radius value"); - const scaling = Math.max(1, norm / r); - return l.div(scaling); - }), - ); + // Compute every norm before allocating any clipped tensor so that a failure + // doesn't leave already clipped layers behind + const norms = await Promise.all(layers.map(frobeniusNorm)); - return new WeightsContainer(clipped); + /** Apply different clipping radius to each layer in the WeightsContainer */ + return new WeightsContainer( + tf.tidy(() => + layers.map((l, i) => l.div(Math.max(1, norms[i] / radius[i]))), + ), + ); } diff --git a/discojs/src/training/trainer.ts b/discojs/src/training/trainer.ts index 34a4c669a..5f8297dc9 100644 --- a/discojs/src/training/trainer.ts +++ b/discojs/src/training/trainer.ts @@ -439,11 +439,14 @@ export class Trainer { const networkWeights = await this.#client.onRoundEndCommunication(roundWeights); - this.model.weights = networkWeights; - // Currently only does something for decentralized clients - // Save weights and cleanup state - this.#client.finishRound(networkWeights); - networkWeights.dispose(); + try { + this.model.weights = networkWeights; + // Currently only does something for decentralized clients + // Save weights and cleanup state + this.#client.finishRound(networkWeights); + } finally { + networkWeights.dispose(); + } return validationDataset !== undefined ? await this.model.evaluate(validationDataset) @@ -488,10 +491,6 @@ async function applyOptimalPrivacy( dpDefaultRadius, ); - const previousEpochWeights = - previous ?? current.map((w) => tf.zerosLike(w)); - const weightsProgress = current.sub(previousEpochWeights); - /** Need to use tighter clipping radius for noise calibration */ const effectiveRadius = "byzantineFaultTolerance" in options @@ -513,17 +512,26 @@ async function applyOptimalPrivacy( sigmaMax: Math.max(...sigmas), }); - const noisyProgress = await privacy.addOptimalNoise( - weightsProgress, - epsilon, - delta, - effectiveRadius, - ); + const previousEpochWeights = + previous ?? current.map((w) => tf.zerosLike(w)); try { - ret = previousEpochWeights.add(noisyProgress); + const weightsProgress = current.sub(previousEpochWeights); + try { + const noisyProgress = await privacy.addOptimalNoise( + weightsProgress, + epsilon, + delta, + effectiveRadius, + ); + try { + ret = previousEpochWeights.add(noisyProgress); + } finally { + noisyProgress.dispose(); + } + } finally { + weightsProgress.dispose(); + } } finally { - weightsProgress.dispose(); - noisyProgress.dispose(); if (previous === undefined) previousEpochWeights.dispose(); } } @@ -532,18 +540,24 @@ async function applyOptimalPrivacy( // might need to change the variable name const previousRoundWeights = previous ?? current.map((w) => tf.zerosLike(w)); - const weightsProgress = current.sub(previousRoundWeights); - const clippedProgress = await privacy.clipNorm( - weightsProgress, - Repeat(options.byzantineFaultTolerance.clippingRadius) - .take(weightsProgress.weights.length) - .toArray(), - ); try { - ret = previousRoundWeights.add(clippedProgress); + const weightsProgress = current.sub(previousRoundWeights); + try { + const clippedProgress = await privacy.clipNorm( + weightsProgress, + Repeat(options.byzantineFaultTolerance.clippingRadius) + .take(weightsProgress.weights.length) + .toArray(), + ); + try { + ret = previousRoundWeights.add(clippedProgress); + } finally { + clippedProgress.dispose(); + } + } finally { + weightsProgress.dispose(); + } } finally { - weightsProgress.dispose(); - clippedProgress.dispose(); if (previous === undefined) previousRoundWeights.dispose(); } } diff --git a/docs/examples/training.ts b/docs/examples/training.ts index a02f0f720..a92a48934 100644 --- a/docs/examples/training.ts +++ b/docs/examples/training.ts @@ -25,11 +25,13 @@ async function runUser( // Create Disco object associated with the server url, the training scheme const disco = new Disco(task, url, { scheme: "federated" }); - // Run training on the dataset - await disco.trainFully(dataset); - - // Disconnect from the remote server - await disco.close(); + try { + // Run training on the dataset + await disco.trainFully(dataset); + } finally { + // Disconnect from the remote server and dispose the model + await disco.close(); + } } type TaskAndDataset = [ diff --git a/webapp/src/components/task_creation_form/TaskCreationForm.vue b/webapp/src/components/task_creation_form/TaskCreationForm.vue index aade49b33..df4bc6f69 100644 --- a/webapp/src/components/task_creation_form/TaskCreationForm.vue +++ b/webapp/src/components/task_creation_form/TaskCreationForm.vue @@ -1144,11 +1144,18 @@ async function onSubmit(form: unknown): Promise { case "image": case "tabular": { const loaded = await tf.loadLayersModel(tf.io.browserFiles([topology])); - loaded.compile({ - loss, - optimizer: tf.train[optimizer.name](optimizer.learningRate), - }); - model = new TFJS(task.dataType, loaded); + try { + loaded.compile({ + loss, + optimizer: tf.train[optimizer.name](optimizer.learningRate), + }); + model = new TFJS(task.dataType, loaded); + } catch (e) { + // compile or TFJS rejected the model, don't leak the loaded weights + loaded.dispose(); + loaded.optimizer?.dispose(); + throw e; + } break; } case "text": From 84d9674b7192d3746d546bce7bf8e5369526c1d0 Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 16:56:18 +0200 Subject: [PATCH 18/23] feat: free tensors if serialization fails --- discojs/src/serialization/weights.spec.ts | 22 +++++++++++++++++++++- discojs/src/serialization/weights.ts | 17 +++++++++++++++-- 2 files changed, 36 insertions(+), 3 deletions(-) diff --git a/discojs/src/serialization/weights.spec.ts b/discojs/src/serialization/weights.spec.ts index f3a0bc998..8f47c1a32 100644 --- a/discojs/src/serialization/weights.spec.ts +++ b/discojs/src/serialization/weights.spec.ts @@ -1,7 +1,8 @@ +import * as tf from "@tensorflow/tfjs"; import { assert, describe, it } from "vitest"; import { WeightsContainer } from "#weights/index"; -import { isEncoded } from "#serialization/coder"; +import { encode as encodeGeneric, isEncoded } from "#serialization/coder"; import { encode, decode } from "#serialization/weights"; describe("weights", () => { @@ -29,4 +30,23 @@ describe("weights", () => { ), ); }); + + const malformedShapes: [string, number[]][] = [ + ["mismatched data length", [3]], + ["negative dimension", [-1]], + ["non-integer dimension", [0.5]], + ["NaN dimension", [Number.NaN]], + ]; + for (const [name, shape] of malformedShapes) + it(`rejects ${name} without leaking tensors`, () => { + const encoded = encodeGeneric([ + { shape: [1], data: new Float32Array([1]) }, + { shape: [2], data: new Float32Array([1, 2]) }, + { shape, data: new Float32Array([1, 2]) }, + ]); + + const before = tf.memory().numTensors; + assert.throws(() => decode(encoded)); + assert.equal(tf.memory().numTensors, before); + }); }); diff --git a/discojs/src/serialization/weights.ts b/discojs/src/serialization/weights.ts index a0441defc..f873673af 100644 --- a/discojs/src/serialization/weights.ts +++ b/discojs/src/serialization/weights.ts @@ -19,11 +19,20 @@ function isSerialized(raw: unknown): raw is Serialized { const { shape, data }: Partial> = raw; if ( - !(Array.isArray(shape) && shape.every((e) => typeof e === "number")) || + !( + Array.isArray(shape) && + shape.every( + (e): e is number => + typeof e === "number" && Number.isSafeInteger(e) && e >= 0, + ) + ) || !(data instanceof Float32Array) ) return false; + // tf.tensor throws if the shape doesn't match the data + if (shape.reduce((acc, e) => acc * e, 1) !== data.length) return false; + const _: Serialized = { shape, data }; return true; @@ -46,5 +55,9 @@ export function decode(encoded: Encoded): WeightsContainer { if (!(Array.isArray(raw) && raw.every(isSerialized))) throw new Error("expected to decode an array of serialized weights"); - return new WeightsContainer(raw.map((w) => tf.tensor(w.data, w.shape))); + // payloads can come from untrusted peers, so free the already built tensors + // if any of them fails + return new WeightsContainer( + tf.tidy(() => raw.map((w) => tf.tensor(w.data, w.shape))), + ); } From 379aab58b80f6d8fca9267e027c1f0ea91b09bd9 Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 17:21:33 +0200 Subject: [PATCH 19/23] fix: leaks in weight equal and reduce & use async .data() instead of .dataSync() --- discojs/src/aggregator/mean.spec.ts | 18 ++--- discojs/src/aggregator/secure.spec.ts | 11 ++-- discojs/src/aggregator/secure_history.spec.ts | 13 ++-- discojs/src/weights/aggregation.spec.ts | 12 ++-- discojs/src/weights/weights_container.spec.ts | 65 +++++++++++++++++++ discojs/src/weights/weights_container.ts | 41 +++++++++--- server/tests/e2e/federated.spec.ts | 10 +-- 7 files changed, 131 insertions(+), 39 deletions(-) create mode 100644 discojs/src/weights/weights_container.spec.ts diff --git a/discojs/src/aggregator/mean.spec.ts b/discojs/src/aggregator/mean.spec.ts index 3bfc367c4..151a672b7 100644 --- a/discojs/src/aggregator/mean.spec.ts +++ b/discojs/src/aggregator/mean.spec.ts @@ -20,8 +20,8 @@ describe("mean aggregator", () => { expect(aggregator.isValidContribution("client 1", 0)).to.be.true; const client1Round0Promise = aggregator.getPromiseForAggregation(); aggregator.add("client 1", WeightsContainer.of([1]), 0); - expect(WeightsContainer.of([1]).equals(await client1Round0Promise)).to.be - .true; + expect(await WeightsContainer.of([1]).equals(await client1Round0Promise)).to + .be.true; expect(aggregator.round).to.equal(1); // round 1 @@ -30,8 +30,8 @@ describe("mean aggregator", () => { aggregator.add("client 1", WeightsContainer.of([1]), 1); const client2Round0Promise = aggregator.getPromiseForAggregation(); aggregator.add("client 2", WeightsContainer.of([2]), 0); - expect(WeightsContainer.of([1.5]).equals(await client2Round0Promise)).to.be - .true; + expect(await WeightsContainer.of([1.5]).equals(await client2Round0Promise)) + .to.be.true; expect(aggregator.round).to.equal(2); // round 2 @@ -42,8 +42,8 @@ describe("mean aggregator", () => { aggregator.add("client 2", WeightsContainer.of([1]), 2); const client3Round2Promise = aggregator.getPromiseForAggregation(); aggregator.add("client 3", WeightsContainer.of([4]), 1); - expect(WeightsContainer.of([2]).equals(await client3Round2Promise)).to.be - .true; + expect(await WeightsContainer.of([2]).equals(await client3Round2Promise)).to + .be.true; expect(aggregator.round).to.equal(3); }); @@ -61,7 +61,7 @@ describe("mean aggregator", () => { aggregator.add(id1, WeightsContainer.of([0], [1]), 0); const result2 = aggregator.getPromiseForAggregation(); aggregator.add(id2, WeightsContainer.of([2], [3]), 0); - expect((await result1).equals(await result2)).to.be.true; + expect(await (await result1).equals(await result2)).to.be.true; expect(await WSIntoArrays(await results)).to.deep.equal([[1], [2]]); }); @@ -103,7 +103,7 @@ describe("mean aggregator", () => { aggregator.registerNode(id2); const result2 = aggregator.getPromiseForAggregation(); aggregator.add(id2, WeightsContainer.of([2], [3]), 0); - expect((await result1).equals(await result2)).to.be.true; + expect(await (await result1).equals(await result2)).to.be.true; expect(aggregator.round).equals(1); // round should be one now }); @@ -143,7 +143,7 @@ describe("mean aggregator", () => { aggregator.registerNode(id2); const result2 = aggregator.getPromiseForAggregation(); aggregator.add(id2, WeightsContainer.of([2], [3]), 0); - expect((await result1).equals(await result2)).to.be.true; + expect(await (await result1).equals(await result2)).to.be.true; expect(aggregator.round).equals(1); }); }); diff --git a/discojs/src/aggregator/secure.spec.ts b/discojs/src/aggregator/secure.spec.ts index 9fe9c88b7..35256c6f0 100644 --- a/discojs/src/aggregator/secure.spec.ts +++ b/discojs/src/aggregator/secure.spec.ts @@ -35,18 +35,19 @@ describe("secret shares test", () => { .toList(); } - it("recover secrets from shares", () => { + it("recover secrets from shares", async () => { const recovered = buildShares().map((shares) => sum(shares)); - assert.isTrue( + const equalities = await Promise.all( ( recovered.zip(secrets) as List<[WeightsContainer, WeightsContainer]> - ).every(([actual, expected]) => actual.equals(expected, epsilon)), + ).map(([actual, expected]) => actual.equals(expected, epsilon)), ); + assert.isTrue(equalities.every((equal) => equal)); }); - it("derive aggregation result from partial sums", () => { + it("derive aggregation result from partial sums", async () => { const actual = avg(buildPartialSums(buildShares())); - assert.isTrue(actual.equals(expected, epsilon)); + assert.isTrue(await actual.equals(expected, epsilon)); }); }); diff --git a/discojs/src/aggregator/secure_history.spec.ts b/discojs/src/aggregator/secure_history.spec.ts index 326f0e720..0399134bc 100644 --- a/discojs/src/aggregator/secure_history.spec.ts +++ b/discojs/src/aggregator/secure_history.spec.ts @@ -28,13 +28,14 @@ describe("Secure history aggregator", function () { }); } - it("recovers secrets from shares", () => { + it("recovers secrets from shares", async () => { const recovered = buildShares().map((shares) => sum(shares)); - assert.isTrue( + const equalities = await Promise.all( ( recovered.zip(secrets) as List<[WeightsContainer, WeightsContainer]> - ).every(([actual, expected]) => actual.equals(expected, epsilon)), + ).map(([actual, expected]) => actual.equals(expected, epsilon)), ); + assert.isTrue(equalities.every((equal) => equal)); }); it("aggregates partial sums with momentum smoothing", async () => { @@ -67,7 +68,7 @@ describe("Secure history aggregator", function () { const expectedSum = sum( sharesRound0.flatMap((x) => x), // flatten to List ); - expect(sumRound0.equals(expectedSum, epsilon)).to.be.true; + expect(await sumRound0.equals(expectedSum, epsilon)).to.be.true; // simulate second communication round partial sums const aggregationPromise2 = aggregator.getPromiseForAggregation(); @@ -80,7 +81,7 @@ describe("Secure history aggregator", function () { // First aggregation with momentum - no previous momentum, so just average const avgPartialSum = avg(partialSums); - expect(sumRound1.equals(avgPartialSum, epsilon)).to.be.true; + expect(await sumRound1.equals(avgPartialSum, epsilon)).to.be.true; // Now we simulate a second round of aggregation with momentum smoothing const dummyPromise = aggregator.getPromiseForAggregation(); @@ -110,7 +111,7 @@ describe("Secure history aggregator", function () { ); // Compare the actual result to the expected smoothed result using momentum - expect(sumRound2.equals(expectedSumRound2, 1e-3)).to.be.true; + expect(await sumRound2.equals(expectedSumRound2, 1e-3)).to.be.true; }); it("behaves similar to SecureAggregator without momentum (beta=0)", async () => { diff --git a/discojs/src/weights/aggregation.spec.ts b/discojs/src/weights/aggregation.spec.ts index cf86e622a..02f6b953d 100644 --- a/discojs/src/weights/aggregation.spec.ts +++ b/discojs/src/weights/aggregation.spec.ts @@ -4,7 +4,7 @@ import { WeightsContainer } from "#weights/weights_container"; import { avg, sum, diff } from "#weights/aggregation"; describe("weights aggregation", () => { - it("avg of weights with two operands", () => { + it("avg of weights with two operands", async () => { const actual = avg([ WeightsContainer.of([1, 2, 3, -1], [-5, 6]), WeightsContainer.of([2, 3, 7, 1], [-10, 5]), @@ -12,7 +12,7 @@ describe("weights aggregation", () => { ]); const expected = WeightsContainer.of([2, 2, 5, 1], [-10, 10]); - assert.isTrue(actual.equals(expected)); + assert.isTrue(await actual.equals(expected)); }); it("avg does not leak intermediate tensors", () => { @@ -65,17 +65,17 @@ describe("weights aggregation", () => { result.dispose(); }); - it("sum of weights with two operands", () => { + it("sum of weights with two operands", async () => { const actual = sum([ [[3, -4], [9]], [[2, 13], [0]], ]); const expected = WeightsContainer.of([5, 9], [9]); - assert.isTrue(actual.equals(expected)); + assert.isTrue(await actual.equals(expected)); }); - it("diff of weights with two operands", () => { + it("diff of weights with two operands", async () => { const actual = diff([ [ [3, -4, 5], @@ -88,6 +88,6 @@ describe("weights aggregation", () => { ]); const expected = WeightsContainer.of([1, -17, 1], [9, 0]); - assert.isTrue(actual.equals(expected)); + assert.isTrue(await actual.equals(expected)); }); }); diff --git a/discojs/src/weights/weights_container.spec.ts b/discojs/src/weights/weights_container.spec.ts new file mode 100644 index 000000000..5dcac8855 --- /dev/null +++ b/discojs/src/weights/weights_container.spec.ts @@ -0,0 +1,65 @@ +import * as tf from "@tensorflow/tfjs"; +import { assert, describe, it } from "vitest"; + +import { WeightsContainer } from "#weights/index"; + +describe("weights container", () => { + it("equals same weights within margin", async () => { + const a = WeightsContainer.of([1, 2], [3]); + const b = WeightsContainer.of([1.05, 2], [3]); + + assert.isTrue(await a.equals(a)); + assert.isFalse(await a.equals(b)); + assert.isTrue(await a.equals(b, 0.1)); + + a.dispose(); + b.dispose(); + }); + + it("doesn't equal containers of different sizes", async () => { + const a = WeightsContainer.of([1], [2]); + const b = WeightsContainer.of([1]); + + assert.isFalse(await a.equals(b)); + assert.isFalse(await b.equals(a)); + + a.dispose(); + b.dispose(); + }); + + it("doesn't equal weights of different shapes", async () => { + // would be equal if broadcasted + const a = WeightsContainer.of([1]); + const b = WeightsContainer.of([1, 1, 1]); + + assert.isFalse(await a.equals(b)); + + a.dispose(); + b.dispose(); + }); + + it("equals does not leak tensors", async () => { + const a = WeightsContainer.of([1, 2], [3]); + const b = WeightsContainer.of([1, 2], [4]); + + const before = tf.memory().numTensors; + await a.equals(b); + await a.equals(a); + assert.equal(tf.memory().numTensors, before); + + a.dispose(); + b.dispose(); + }); + + it("reduce only keeps the result", async () => { + const weights = WeightsContainer.of([1], [2], [3], [4]); + + const before = tf.memory().numTensors; + const reduced = weights.reduce((acc, t) => acc.add(t)); + assert.equal(tf.memory().numTensors, before + 1); + assert.deepEqual(Array.from(await reduced.data()), [10]); + + reduced.dispose(); + weights.dispose(); + }); +}); diff --git a/discojs/src/weights/weights_container.ts b/discojs/src/weights/weights_container.ts index 633fe4260..dbe07cfec 100644 --- a/discojs/src/weights/weights_container.ts +++ b/discojs/src/weights/weights_container.ts @@ -84,8 +84,12 @@ export class WeightsContainer { return new WeightsContainer(this._weights.map(fn)); } + /** + * Folds the weights with the given binary operator. + * Intermediate accumulators are disposed, only the final one is kept. + */ reduce(fn: (acc: tf.Tensor, t: tf.Tensor) => tf.Tensor): tf.Tensor { - return this._weights.reduce(fn); + return tf.tidy(() => this._weights.reduce(fn)); } /** @@ -101,13 +105,34 @@ export class WeightsContainer { return WeightsContainer.of(...this.weights, ...other.weights); } - equals(other: WeightsContainer, margin = 0): boolean { - return this._weights - .zip(other._weights) - .every( - ([w1, w2]) => - w1.sub(w2).abs().lessEqual(margin).all().dataSync()[0] === 1, - ); + /** + * Checks that both containers hold weights of the same shapes, entry-wise + * equal up to the given margin. + */ + async equals(other: WeightsContainer, margin = 0): Promise { + if (this._weights.size !== other._weights.size) return false; + const pairs = this._weights.zip(other._weights) as List< + [tf.Tensor, tf.Tensor] + >; + // otherwise sub would broadcast + if (!pairs.every(([w1, w2]) => tf.util.arraysEqual(w1.shape, w2.shape))) + return false; + if (pairs.isEmpty()) return true; + + const allClose = tf.tidy(() => + tf + .stack( + pairs + .map(([w1, w2]) => w1.sub(w2).abs().lessEqual(margin).all()) + .toArray(), + ) + .all(), + ); + try { + return (await allClose.data())[0] === 1; + } finally { + allClose.dispose(); + } } dispose(): void { diff --git a/server/tests/e2e/federated.spec.ts b/server/tests/e2e/federated.spec.ts index 708b8ca96..cab5a2973 100644 --- a/server/tests/e2e/federated.spec.ts +++ b/server/tests/e2e/federated.spec.ts @@ -115,7 +115,7 @@ describe("end-to-end federated", () => { expect(lastEpoch.training.accuracy).to.be.greaterThan(0.4); expect(lastEpoch.validation?.accuracy).to.be.greaterThan(0.4); } - assert.isTrue(m1.equals(m2) && m2.equals(m3)); + assert.isTrue((await m1.equals(m2)) && (await m2.equals(m3))); }); it("two titanic users reach consensus", { timeout: 50_000 }, async () => { @@ -143,7 +143,7 @@ describe("end-to-end federated", () => { expect(lastEpoch.training.accuracy).to.be.greaterThan(0.4); expect(lastEpoch.validation?.accuracy).to.be.greaterThan(0.4); } - assert.isTrue(m1.equals(m2)); + assert.isTrue(await m1.equals(m2)); }); it("two lus_covid users reach consensus", { timeout: 200_000 }, async () => { @@ -170,7 +170,7 @@ describe("end-to-end federated", () => { expect(lastEpoch.training.accuracy).to.be.greaterThan(0.4); expect(lastEpoch.validation?.accuracy).to.be.greaterThan(0.4); } - assert.isTrue(m1.equals(m2)); + assert.isTrue(await m1.equals(m2)); }); it("two wikitext reach consensus", { timeout: 500_000 }, async () => { @@ -202,7 +202,7 @@ describe("end-to-end federated", () => { runUser(url, task, dataset, false), runUser(url, task, dataset, false), ]); - assert.isTrue(r1[0].equals(r2[0])); + assert.isTrue(await r1[0].equals(r2[0])); }); /** @@ -485,7 +485,7 @@ describe("end-to-end federated", () => { expect(lastEpoch.training.accuracy).to.be.greaterThan(0.4); expect(lastEpoch.validation?.accuracy).to.be.greaterThan(0.4); } - assert.isTrue(m1.equals(m2) && m2.equals(m3)); + assert.isTrue((await m1.equals(m2)) && (await m2.equals(m3))); }, ); From c56dec1605df03e248ca21316dfb5eba9f9b07d0 Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Mon, 28 Sep 2026 17:35:00 +0200 Subject: [PATCH 20/23] fix: dispose model in cli scripts --- cli/src/benchmark_gpt.ts | 4 ++-- cli/src/evaluate_finetuned_gpt2.ts | 2 +- cli/src/hellaswag_gpt.ts | 6 +++++- cli/src/measure_memorization_gpt2.ts | 2 +- cli/src/train_gpt.ts | 2 +- onnx-converter/src/convert_onnx.ts | 10 +++++++--- 6 files changed, 17 insertions(+), 9 deletions(-) diff --git a/cli/src/benchmark_gpt.ts b/cli/src/benchmark_gpt.ts index 4e206317e..6bdc64ceb 100644 --- a/cli/src/benchmark_gpt.ts +++ b/cli/src/benchmark_gpt.ts @@ -130,7 +130,7 @@ async function main(args: Required): Promise { .batch(batchSize); // Init and train the model - const model = new GPT(config); + using model = new GPT(config); console.log( `\tmodel type ${modelType} \n\tbatch size ${batchSize} \n\tcontext length ${contextLength}`, ); @@ -151,7 +151,7 @@ async function main(args: Required): Promise { * Inference benchmark */ } else { - const model = await loadModelFromDisk(modelPath); + using model = await loadModelFromDisk(modelPath); if (!(model instanceof GPT)) { throw new Error("Loaded model isn't a GPT model"); } diff --git a/cli/src/evaluate_finetuned_gpt2.ts b/cli/src/evaluate_finetuned_gpt2.ts index c17d291dd..b4b73aacb 100644 --- a/cli/src/evaluate_finetuned_gpt2.ts +++ b/cli/src/evaluate_finetuned_gpt2.ts @@ -493,7 +493,7 @@ async function main() { const tokenizer = await Tokenizer.from_pretrained("Xenova/gpt2"); console.log("Loading model..."); - const model = await loadModelFromDisk(args.modelPath); + using model = await loadModelFromDisk(args.modelPath); if (!(model instanceof GPT)) { throw new Error("Model must be GPT"); diff --git a/cli/src/hellaswag_gpt.ts b/cli/src/hellaswag_gpt.ts index c68862ddb..17d06d310 100644 --- a/cli/src/hellaswag_gpt.ts +++ b/cli/src/hellaswag_gpt.ts @@ -115,7 +115,11 @@ async function main(): Promise { model = (await modelDecode(encodedModel)) as GPT; break; } - await evaluateModel(model, args.numDataPoints); + try { + await evaluateModel(model, args.numDataPoints); + } finally { + model.dispose(); + } console.log("Benchmark completed!"); } diff --git a/cli/src/measure_memorization_gpt2.ts b/cli/src/measure_memorization_gpt2.ts index c3dc31415..66e7162ad 100644 --- a/cli/src/measure_memorization_gpt2.ts +++ b/cli/src/measure_memorization_gpt2.ts @@ -343,7 +343,7 @@ async function main() { const tokenizer = await Tokenizer.from_pretrained("Xenova/gpt2"); console.log("Loading model..."); - const loadedModel = await loadModelFromDisk(args.modelPath); + using loadedModel = await loadModelFromDisk(args.modelPath); if (!(loadedModel instanceof GPT)) { throw new Error("modelPath must point to a Disco GPT model"); } diff --git a/cli/src/train_gpt.ts b/cli/src/train_gpt.ts index f478db14f..821ce6d82 100644 --- a/cli/src/train_gpt.ts +++ b/cli/src/train_gpt.ts @@ -27,7 +27,7 @@ async function main(): Promise { .repeat() .batch(8); - const model = new GPT(config); + using model = new GPT(config); for await (const logs of model.train(tokenDataset, undefined)) { console.log(logs); } diff --git a/onnx-converter/src/convert_onnx.ts b/onnx-converter/src/convert_onnx.ts index 39a6b98e8..a52ee7149 100644 --- a/onnx-converter/src/convert_onnx.ts +++ b/onnx-converter/src/convert_onnx.ts @@ -33,7 +33,7 @@ async function main() { console.log("ONNX model loaded successfully"); // Init empty TF.js model - const gptModel = new GPT({ + using gptModel = new GPT({ modelType: "gpt2", contextLength: GPT2_CONTEXT_LENGTH, }); @@ -62,7 +62,7 @@ async function main() { throw new Error(`Undefined layer dimensions for ${tensor.name}`); const dims = tensor.dims.map((d) => Number(d)); const flatData = parseTensorData(tensor); - let tfTensor = tf.tensor(flatData).reshape(dims); + let tfTensor = tf.tensor(flatData, dims); if (tensor.name === "transformer.wpe.weight") { if (dims.length !== 2) throw new Error( @@ -72,7 +72,9 @@ async function main() { throw new Error( `ONNX positional embeddings only support context length ${dims[0]}, requested ${GPT2_CONTEXT_LENGTH}.`, ); - tfTensor = tfTensor.slice([0, 0], [GPT2_CONTEXT_LENGTH, dims[1]]); + const full = tfTensor; + tfTensor = full.slice([0, 0], [GPT2_CONTEXT_LENGTH, dims[1]]); + full.dispose(); } preTrainedWeights = preTrainedWeights.set(tfjsName, tfTensor); } @@ -95,6 +97,8 @@ async function main() { }); gptLayersModel.setWeights(finalWeights); // shape or transpose mismatch will throw here + // the model's variables now hold their own reference to the data + preTrainedWeights.forEach((t) => t.dispose()); const encoded = await modelEncode(gptModel); await fsPromise.writeFile(OUTPUT_FILENAME, encoded); From 4544abfb1c759448260a5b968824306395e6b70e Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Tue, 29 Sep 2026 13:17:53 +0200 Subject: [PATCH 21/23] fix: free CIFAR10 model unused layers --- .../CIFAR10ClassifierModel.spec.ts | 15 +++++++++++++++ .../implementations/CIFAR10ClassifierModel.ts | 7 +++++++ 2 files changed, 22 insertions(+) create mode 100644 discojs/src/models/implementations/CIFAR10ClassifierModel.spec.ts diff --git a/discojs/src/models/implementations/CIFAR10ClassifierModel.spec.ts b/discojs/src/models/implementations/CIFAR10ClassifierModel.spec.ts new file mode 100644 index 000000000..6cdae8cd1 --- /dev/null +++ b/discojs/src/models/implementations/CIFAR10ClassifierModel.spec.ts @@ -0,0 +1,15 @@ +import * as tf from "@tensorflow/tfjs"; +import { describe, expect, it } from "vitest"; + +import { getModel } from "#models/implementations/CIFAR10ClassifierModel"; + +describe("CIFAR10 classifier model", () => { + it("frees every tensor on dispose", async () => { + const before = tf.memory().numTensors; + + const model = await getModel(); + model.dispose(); + + expect(tf.memory().numTensors).to.equal(before); + }); +}); diff --git a/discojs/src/models/implementations/CIFAR10ClassifierModel.ts b/discojs/src/models/implementations/CIFAR10ClassifierModel.ts index 0435ce453..da627c0ae 100644 --- a/discojs/src/models/implementations/CIFAR10ClassifierModel.ts +++ b/discojs/src/models/implementations/CIFAR10ClassifierModel.ts @@ -20,6 +20,13 @@ export async function getModel() { name: "modelModified", }); + // free the weights of the original classification head (conv_preds, ...) + // which are not part of the new model + const kept = new Set(model.layers); + mobilenet.layers + .filter((layer) => !kept.has(layer)) + .forEach((layer) => layer.dispose()); + model.compile({ optimizer: "sgd", loss: "categoricalCrossentropy", From bb2c11608989c2b8e37c15f847b2ddd6bec0044b Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Tue, 29 Sep 2026 14:35:00 +0200 Subject: [PATCH 22/23] doc: update tensor management conventions --- docs/DISCOJS.md | 32 ++++++++++++++++++++++++++++---- 1 file changed, 28 insertions(+), 4 deletions(-) diff --git a/docs/DISCOJS.md b/docs/DISCOJS.md index a66f9d815..33ff22847 100644 --- a/docs/DISCOJS.md +++ b/docs/DISCOJS.md @@ -204,8 +204,32 @@ async function frobeniusNorm(weights: tf.Tensor): Promise { #### Ownership conventions -To know who is responsible for disposing a tensor, we follow these conventions: +Every tensor has exactly one **owner**, who is responsible for disposing it. Code that uses a tensor without owning it only **borrows** it: it must not dispose it, and must not use it after the owner may have disposed it. -1. A function does not dispose its inputs, the caller remains responsible for them. -2. A function returning tensors (or a `WeightsContainer`) transfers their ownership to the caller, who has to dispose them. -3. Returned tensors must not alias the inputs (e.g., return `input.clone()` rather than `input`), otherwise disposing one would dispose the other. +| Situation | Who owns the tensor | +| -------------------------------------------------------------------------------------- | ----------------------------------------------------------------- | +| Arguments passed to a function or method | The caller: the callee borrows them and must not dispose them | +| Tensors returned by a function or method | The caller, who has to dispose them | +| Tensors returned by a getter (`model.weights`, `weightsContainer.weights`, `.get(i)`) | The object: the caller borrows them and must not dispose them | +| Tensors passed to a constructor or wrapper (`new WeightsContainer(tensors)`) | The new object, which disposes them in its own `dispose()` | +| Tensors kept by an object beyond a call (stored in a field, a map, an aggregator, ...) | The object, which must store a `clone()` rather than the argument | +| Tensors emitted in an event (e.g. the aggregator's `"aggregation"` event) | The listener that consumes them: there must be exactly one | +| Objects handed over explicitly (e.g. `trainer.releaseModel()`) | The caller, the previous owner forgets its reference | + +#### Using memory efficiently + +In the browser, a tab can only use a limited amount of memory, and models such as GPT are large. What matters is not only leaks but also **peak memory**, i.e., how many tensors are alive at the same time. + +- **`clone()` does not copy data.** A clone shares its buffer with the original, and the buffer is freed once the last tensor referencing it is disposed. Cloning to respect the rules above is therefore cheap. +- **A clone keeps old values alive.** Assigning new values to a variable (e.g. setting `model.weights`) makes it point to a new buffer, while existing clones keep the old one. A snapshot of the model weights thus costs a full copy of the model as soon as the model is updated. Dispose snapshots as soon as they are not needed anymore. +- **Setting weights does not copy them either.** After `model.weights = weights`, the model and `weights` share the same buffers: the caller still owns `weights` and can dispose it right away, without affecting the model. +- **Keep `tf.tidy` scopes short.** `tf.tidy` only disposes intermediate tensors when `fn` returns, so wrapping a whole loop keeps every iteration's intermediates alive until the end. In loops, use one `tf.tidy` per iteration and dispose the accumulator manually: + +```ts +let acc = tf.zeros([10]); +for (const t of tensors) { + const next = tf.tidy(() => acc.add(t.square())); + acc.dispose(); + acc = next; +} +``` From 7b599acf44e906ea415176271ecf1ae445c3e69c Mon Sep 17 00:00:00 2001 From: Julien Vignoud Date: Wed, 30 Sep 2026 15:04:20 +0200 Subject: [PATCH 23/23] fix: weight container methods don't alias input --- discojs/src/weights/weights_container.spec.ts | 72 +++++++++++++++++++ discojs/src/weights/weights_container.ts | 32 +++++++-- 2 files changed, 98 insertions(+), 6 deletions(-) diff --git a/discojs/src/weights/weights_container.spec.ts b/discojs/src/weights/weights_container.spec.ts index 5dcac8855..263f0f46d 100644 --- a/discojs/src/weights/weights_container.spec.ts +++ b/discojs/src/weights/weights_container.spec.ts @@ -62,4 +62,76 @@ describe("weights container", () => { reduced.dispose(); weights.dispose(); }); + + it("map doesn't alias the weights", async () => { + const weights = WeightsContainer.of([1], [2]); + const expected = WeightsContainer.of([1], [2]); + + const before = tf.memory().numTensors; + const mapped = weights.map((t) => t); + mapped.dispose(); + assert.equal(tf.memory().numTensors, before); + + assert.isFalse(weights.weights.some((t) => t.isDisposed)); + assert.isTrue(await weights.equals(expected)); + + weights.dispose(); + expected.dispose(); + }); + + it("mapWith doesn't alias the weights", () => { + const a = WeightsContainer.of([1], [2]); + const b = WeightsContainer.of([3], [4]); + + const before = tf.memory().numTensors; + a.mapWith(b, (w) => w).dispose(); + a.mapWith(b, (_, w) => w).dispose(); + assert.equal(tf.memory().numTensors, before); + + assert.isFalse(a.weights.some((t) => t.isDisposed)); + assert.isFalse(b.weights.some((t) => t.isDisposed)); + + a.dispose(); + b.dispose(); + }); + + it("reduce doesn't alias the weights", async () => { + const single = WeightsContainer.of([1]); + const weights = WeightsContainer.of([1], [2]); + + const before = tf.memory().numTensors; + const reducedSingle = single.reduce((acc, t) => acc.add(t)); + const reducedFirst = weights.reduce((acc) => acc); + assert.equal(tf.memory().numTensors, before + 2); + assert.deepEqual(Array.from(await reducedSingle.data()), [1]); + assert.deepEqual(Array.from(await reducedFirst.data()), [1]); + + reducedSingle.dispose(); + reducedFirst.dispose(); + assert.equal(tf.memory().numTensors, before); + assert.isFalse(single.weights.some((t) => t.isDisposed)); + assert.isFalse(weights.weights.some((t) => t.isDisposed)); + + single.dispose(); + weights.dispose(); + }); + + it("concat doesn't alias the weights", async () => { + const a = WeightsContainer.of([1]); + const b = WeightsContainer.of([2], [3]); + const expected = WeightsContainer.of([1], [2], [3]); + + const before = tf.memory().numTensors; + const concatenated = a.concat(b); + assert.isTrue(await concatenated.equals(expected)); + concatenated.dispose(); + assert.equal(tf.memory().numTensors, before); + + assert.isFalse(a.weights.some((t) => t.isDisposed)); + assert.isFalse(b.weights.some((t) => t.isDisposed)); + + a.dispose(); + b.dispose(); + expected.dispose(); + }); }); diff --git a/discojs/src/weights/weights_container.ts b/discojs/src/weights/weights_container.ts index dbe07cfec..7d2fbdc36 100644 --- a/discojs/src/weights/weights_container.ts +++ b/discojs/src/weights/weights_container.ts @@ -68,12 +68,15 @@ export class WeightsContainer { fn: (a: tf.Tensor, b: tf.Tensor) => tf.Tensor, ): WeightsContainer { return new WeightsContainer( - this._weights - .zip(other._weights) - .map(([w1, w2]) => fn(w1, w2 as tf.Tensor)), + this._weights.zip(other._weights).map(([w1, w2]) => { + const mapped = fn(w1, w2 as tf.Tensor); + // `fn` may return one of its inputs, in which case we clone it + return mapped === w1 || mapped === w2 ? mapped.clone() : mapped; + }), ); } + // The result never aliases the container's weights. map(fn: (t: tf.Tensor, i: number) => tf.Tensor): WeightsContainer; map(fn: (t: tf.Tensor) => tf.Tensor): WeightsContainer; map( @@ -81,15 +84,26 @@ export class WeightsContainer { | ((t: tf.Tensor) => tf.Tensor) | ((t: tf.Tensor, i: number) => tf.Tensor), ): WeightsContainer { - return new WeightsContainer(this._weights.map(fn)); + return new WeightsContainer( + this._weights.map((t, i) => { + const mapped = fn(t, i); + // `fn` may return its input (e.g. `map((t) => t)`) in which case we clone it + return mapped === t ? t.clone() : mapped; + }), + ); } /** * Folds the weights with the given binary operator. * Intermediate accumulators are disposed, only the final one is kept. + * The result never aliases the container's weights. */ reduce(fn: (acc: tf.Tensor, t: tf.Tensor) => tf.Tensor): tf.Tensor { - return tf.tidy(() => this._weights.reduce(fn)); + return tf.tidy(() => { + const reduced = this._weights.reduce(fn); + // a single weight is returned as is by `List.reduce`, and `fn` may return one of its inputs + return this._weights.includes(reduced) ? reduced.clone() : reduced; + }); } /** @@ -101,8 +115,14 @@ export class WeightsContainer { return this._weights.get(index); } + /** + * Concatenates this weights container with another one. + * @returns A new weights container holding clones of both containers' weights + */ concat(other: WeightsContainer): WeightsContainer { - return WeightsContainer.of(...this.weights, ...other.weights); + return new WeightsContainer( + this._weights.concat(other._weights).map((t) => t.clone()), + ); } /**