Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 commits
Select commit Hold shift + click to select a range
171b7dd
fix: byzantine aggregator tensor leaks
JulienVig Sep 25, 2026
48ac018
fix: dispose aggregators on close
JulienVig Sep 25, 2026
0c126a3
fix: percentile_clipping aggregator tensor leaks
JulienVig Sep 25, 2026
ae52bdc
fix: secure_history aggregator tensor leaks
JulienVig Sep 25, 2026
31254de
fix: aggregators makePayloads return copies
JulienVig Sep 25, 2026
5a90f2d
fix: dispose decoded GPT-2 weights
JulienVig Sep 28, 2026
fc66bc6
fix: dispose attention mask and token embeddings
JulienVig Sep 28, 2026
5abb024
fix: dispose model after serialization
JulienVig Sep 28, 2026
0f110d1
fix: modelSync leak
JulienVig Sep 28, 2026
763e8a4
fix: dispose lossTensor
JulienVig Sep 28, 2026
ea8cee9
fix: avoid calling already disposed model
JulienVig Sep 28, 2026
9ed5912
doc: use trainer.releaseModel
JulienVig Sep 28, 2026
7fe98bc
fix: stop inference when component unmounts
JulienVig Sep 28, 2026
e3831aa
fix: dispose model after creating new task
JulienVig Sep 28, 2026
aa93f52
fix: dispose llm when leaving ChatUI
JulienVig Sep 28, 2026
3f86b20
fix: dispose unused models
JulienVig Sep 28, 2026
2272d9a
fix: dispose on error
JulienVig Sep 28, 2026
f22f16a
feat: free tensors if serialization fails
JulienVig Sep 28, 2026
9ae9329
fix: leaks in weight equal and reduce & use async .data() instead of …
JulienVig Sep 28, 2026
d9374e0
fix: dispose model in cli scripts
JulienVig Sep 28, 2026
4cd4a62
fix: free CIFAR10 model unused layers
JulienVig Sep 29, 2026
3870e20
doc: update tensor management conventions
JulienVig Sep 29, 2026
10cd762
fix: weight container methods don't alias input
JulienVig Sep 30, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 19 additions & 14 deletions cli/src/args.ts
Original file line number Diff line number Diff line change
Expand Up @@ -444,21 +444,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;
Expand Down
4 changes: 2 additions & 2 deletions cli/src/benchmark_gpt.ts
Original file line number Diff line number Diff line change
Expand Up @@ -130,7 +130,7 @@ async function main(args: Required<CLIArguments>): Promise<void> {
.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}`,
);
Expand All @@ -151,7 +151,7 @@ async function main(args: Required<CLIArguments>): Promise<void> {
* 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");
}
Expand Down
28 changes: 14 additions & 14 deletions cli/src/evaluate_finetuned_gpt2.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down Expand Up @@ -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[] = [];
Expand All @@ -275,7 +270,6 @@ async function scoreContinuations(
});

if (targetIndexes.length === 0) {
inputTensor.dispose();
return scoredInputs.map((scoredInput) => ({
score: Number.NEGATIVE_INFINITY,
promptTokens: promptTokens.length,
Expand All @@ -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],
Expand All @@ -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[];

Expand All @@ -321,10 +325,6 @@ async function scoreContinuations(
usedInputTokens: scoredInput.truncatedInputTokens.length,
}));

inputTensor.dispose();
logits.dispose();
targetLogProbs.dispose();

return results;
}

Expand Down Expand Up @@ -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");
Expand Down
6 changes: 5 additions & 1 deletion cli/src/hellaswag_gpt.ts
Original file line number Diff line number Diff line change
Expand Up @@ -115,7 +115,11 @@ async function main(): Promise<void> {
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!");
}
Expand Down
26 changes: 10 additions & 16 deletions cli/src/measure_memorization_gpt2.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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<tf.Tensor1D>().div(temperature);
const { values: topKLogits, indices: topKTokens } = tf.topk(scaled, topK);
Expand All @@ -213,12 +207,12 @@ async function sampleGenerateGPT2(
return topKTokens.gather(sampledIndex).squeeze<tf.Scalar>();
});

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);
}
Expand Down Expand Up @@ -349,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");
}
Expand Down
2 changes: 1 addition & 1 deletion cli/src/train_gpt.ts
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@ async function main(): Promise<void> {
.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);
}
Expand Down
19 changes: 15 additions & 4 deletions discojs/src/aggregator.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -122,11 +129,15 @@ export async function communicate<A extends Aggregator>(
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;
Expand Down
1 change: 1 addition & 0 deletions discojs/src/aggregator/aggregator.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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<NodeID, WeightsContainer>;
Expand Down
47 changes: 47 additions & 0 deletions discojs/src/aggregator/byzantine.spec.ts
Original file line number Diff line number Diff line change
@@ -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";
Expand Down Expand Up @@ -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);
});
});
Loading
Loading