feat: trace DeepSeek Chat completion depth

This commit is contained in:
wuyang
2026-07-30 00:24:27 +08:00
parent fb44d15bb8
commit 8bb488f275
23 changed files with 1507052 additions and 33 deletions
@@ -0,0 +1,461 @@
import { createHash } from "node:crypto";
import { readFileSync, statSync, writeFileSync } from "node:fs";
import { resolve } from "node:path";
const root = resolve(import.meta.dirname, "..");
const paths = {
baseline: resolve(root, "src/data/deepseek-v2-lite-chat-behavior.json"),
completion: resolve(
root,
"src/data/deepseek-v2-lite-chat-completion-512.json",
),
completionEval: resolve(
root,
"src/data/deepseek-v2-lite-chat-completion-512-eval.json",
),
completionRepro: resolve(
root,
"src/data/deepseek-v2-lite-chat-completion-512-repro-1pd.json",
),
depth: resolve(
root,
"src/data/deepseek-v2-lite-chat-full-depth.json",
),
depthRepro: resolve(
root,
"src/data/deepseek-v2-lite-chat-full-depth-repro-1pd.json",
),
output: resolve(
root,
"src/data/deepseek-v2-lite-chat-completion-depth-compact.json",
),
};
const readJson = (path) => JSON.parse(readFileSync(path, "utf8"));
const sha256 = (path) => createHash("sha256")
.update(readFileSync(path))
.digest("hex");
const artifact = (path) => ({
bytes: statSync(path).size,
sha256: sha256(path),
});
const baseline = readJson(paths.baseline);
const completion = readJson(paths.completion);
const completionEval = readJson(paths.completionEval);
const completionRepro = readJson(paths.completionRepro);
const depth = readJson(paths.depth);
const depthRepro = readJson(paths.depthRepro);
const completionArtifact = artifact(paths.completion);
const baselineArtifact = artifact(paths.baseline);
if (
completionEval.input.behavior_sha256 !== completionArtifact.sha256
|| completionEval.input.baseline_sha256 !== baselineArtifact.sha256
) {
throw new Error("completion evaluator input hash contract failed");
}
const outputKey = (sourceId, condition) => `${sourceId}\0${condition}`;
const completionByKey = new Map(
completion.sources.flatMap((source) => source.outputs.map((output) => [
outputKey(source.id, output.condition),
output,
])),
);
const completionReproduction = {
sources: completionRepro.sources.length,
cells: 0,
promptHashExact: 0,
generatedTokenIdsExact: 0,
generatedTextExact: 0,
eosStateExact: 0,
truncationStateExact: 0,
};
for (const source of completionRepro.sources) {
for (const output of source.outputs) {
const formal = completionByKey.get(outputKey(source.id, output.condition));
if (!formal) {
throw new Error(
`completion formal output missing: ${source.id}/${output.condition}`,
);
}
completionReproduction.cells += 1;
completionReproduction.promptHashExact += (
formal.prompt_token_ids_sha256 === output.prompt_token_ids_sha256
);
completionReproduction.generatedTokenIdsExact += (
JSON.stringify(formal.generated_token_ids)
=== JSON.stringify(output.generated_token_ids)
);
completionReproduction.generatedTextExact += formal.text === output.text;
completionReproduction.eosStateExact += formal.hit_eos === output.hit_eos;
completionReproduction.truncationStateExact += (
formal.stopped_at_max_new_tokens
=== output.stopped_at_max_new_tokens
);
}
}
for (const [name, value] of Object.entries(completionReproduction)) {
if (
!["sources", "cells"].includes(name)
&& value !== completionReproduction.cells
) {
throw new Error(
`completion reproduction mismatch: ${name}=${value}`,
);
}
}
const stableDepthSource = (source) => {
const copy = structuredClone(source);
delete copy.forward_seconds;
delete copy.peak_cuda_memory_allocated_bytes;
return copy;
};
const depthSourceById = new Map(
depth.sources.map((source) => [source.id, source]),
);
const depthReproduction = {
sources: depthRepro.sources.length,
sourceObjectsExactAfterRuntimeStrip: 0,
hiddenStagesPerSource: 29,
routerLayersPerSource: 26,
hiddenTensorHashesExact: 0,
hiddenTensorHashesCompared: 0,
orderedRouteHashesExact: 0,
orderedRouteHashesCompared: 0,
routeWeightHashesExact: 0,
routeWeightHashesCompared: 0,
};
for (const rerunSource of depthRepro.sources) {
const formalSource = depthSourceById.get(rerunSource.id);
if (!formalSource) {
throw new Error(`depth formal source missing: ${rerunSource.id}`);
}
depthReproduction.sourceObjectsExactAfterRuntimeStrip += (
JSON.stringify(stableDepthSource(formalSource))
=== JSON.stringify(stableDepthSource(rerunSource))
);
for (const [stage, stageValue] of Object.entries(
rerunSource.hidden_stages,
)) {
for (const [condition, conditionValue] of Object.entries(
stageValue.conditions,
)) {
for (const scope of ["target_content", "full_input"]) {
depthReproduction.hiddenTensorHashesCompared += 1;
depthReproduction.hiddenTensorHashesExact += (
formalSource.hidden_stages[stage].conditions[condition][scope]
.tensor_sha256
=== conditionValue[scope].tensor_sha256
);
}
}
}
for (const [layer, layerValue] of Object.entries(
rerunSource.router_layers,
)) {
for (const [condition, conditionValue] of Object.entries(
layerValue.conditions,
)) {
for (const scope of ["target_content", "full_input"]) {
const formalScope = (
formalSource.router_layers[layer].conditions[condition][scope]
);
depthReproduction.orderedRouteHashesCompared += 1;
depthReproduction.orderedRouteHashesExact += (
formalScope.ordered_route_sha256
=== conditionValue[scope].ordered_route_sha256
);
depthReproduction.routeWeightHashesCompared += 1;
depthReproduction.routeWeightHashesExact += (
formalScope.route_weight_sha256
=== conditionValue[scope].route_weight_sha256
);
}
}
}
}
if (
depthReproduction.sourceObjectsExactAfterRuntimeStrip
!== depthReproduction.sources
|| depthReproduction.hiddenTensorHashesExact
!== depthReproduction.hiddenTensorHashesCompared
|| depthReproduction.orderedRouteHashesExact
!== depthReproduction.orderedRouteHashesCompared
|| depthReproduction.routeWeightHashesExact
!== depthReproduction.routeWeightHashesCompared
) {
throw new Error(
`depth reproduction mismatch: ${JSON.stringify(depthReproduction)}`,
);
}
const conditions = Object.keys(completion.generation_contract.conditions);
const edgeOrder = Object.keys(depth.trace_contract.edge_pairs);
const edgeLabels = {
system_eos: "System on − off · EOS",
system_bos: "System on − off · BOS",
system_x: "System on − off · x",
system_period: "System on − off · 句点",
bos_at_s0: "BOS − EOS · system off",
bos_at_s1: "BOS − EOS · system on",
x_at_s0: "x − EOS · system off",
x_at_s1: "x − EOS · system on",
period_at_s0: "句点 − EOS · system off",
period_at_s1: "句点 − EOS · system on",
};
const domains = ["all", "english", "chinese", "code", "math"];
const hiddenStages = [
"embedding",
...Array.from({ length: 27 }, (_, index) => (
`layer_${String(index).padStart(2, "0")}`
)),
"final_norm",
];
const routerLayers = Array.from({ length: 26 }, (_, index) => (
`layer_${String(index + 1).padStart(2, "0")}`
));
const sourcesForDomain = (domain) => (
domain === "all"
? depth.sources
: depth.sources.filter((source) => source.domain === domain)
);
const mean = (values) => (
values.reduce((sum, value) => sum + value, 0) / values.length
);
const aggregateHidden = (sources, stage, edge) => {
const rows = sources.map(
(source) => (
source.hidden_stages[stage].target_comparisons[edge]
),
);
const tokens = rows.reduce((sum, row) => sum + row.tokens, 0);
return {
tokens,
exactRows: rows.reduce(
(sum, row) => sum + row.exact_hidden_rows,
0,
),
meanCosine: rows.reduce(
(sum, row) => sum + row.mean_cosine_similarity * row.tokens,
0,
) / tokens,
meanRelativeL2: mean(rows.map((row) => row.mean_relative_l2)),
maxAbsDelta: Math.max(...rows.map((row) => row.max_abs_delta)),
};
};
const aggregateRouter = (sources, layer, edge) => {
const rows = sources.map(
(source) => (
source.router_layers[layer].target_comparisons[edge]
),
);
const tokens = rows.reduce((sum, row) => sum + row.tokens, 0);
const orderedExact = rows.reduce(
(sum, row) => sum + row.ordered_topk_exact_tokens,
0,
);
const setExact = rows.reduce(
(sum, row) => sum + row.set_exact_tokens,
0,
);
return {
tokens,
orderedExact,
orderedExactRate: orderedExact / tokens,
setExact,
setExactRate: setExact / tokens,
meanSetJaccard: rows.reduce(
(sum, row) => sum + row.mean_set_jaccard * row.tokens,
0,
) / tokens,
meanTokenWeightedTv: mean(
rows.map((row) => row.mean_token_weighted_tv),
),
meanAggregateLoadTv: mean(
rows.map((row) => row.aggregate_load_tv),
),
};
};
const hiddenSeries = Object.fromEntries(domains.map((domain) => {
const sources = sourcesForDomain(domain);
return [domain, Object.fromEntries(edgeOrder.map((edge) => [
edge,
hiddenStages.map((stage) => ({
stage,
...aggregateHidden(sources, stage, edge),
})),
]))];
}));
const routerSeries = Object.fromEntries(domains.map((domain) => {
const sources = sourcesForDomain(domain);
return [domain, Object.fromEntries(edgeOrder.map((edge) => [
edge,
routerLayers.map((layer) => ({
layer,
...aggregateRouter(sources, layer, edge),
})),
]))];
}));
const evaluationRows = completionEval.rows.map((row) => ({
sourceId: row.source_id,
domain: row.domain,
condition: row.condition,
generatedTokens: row.generated_tokens,
hitEos: row.hit_eos,
completionClass: row.completion_class,
taskEvaluation: row.task_evaluation,
}));
const incompleteRows = evaluationRows.filter((row) => !row.hitEos);
const codeRows = evaluationRows.filter((row) => row.domain === "code");
const mathRows = evaluationRows.filter((row) => row.domain === "math");
const codeStatuses = Object.groupBy(
codeRows,
(row) => row.taskEvaluation.execution?.status ?? "not_run",
);
const mathExtraction = Object.groupBy(
mathRows,
(row) => row.taskEvaluation.extraction_method,
);
const baselineNaturalEos = baseline.sources.reduce(
(sum, source) => (
sum + source.outputs.filter((output) => output.hit_eos).length
),
0,
);
const boundaryCrossingTokens = depth.sources.reduce(
(sum, source) => (
sum + Object.values(source.boundary_crossing_tokens_by_condition)
.reduce((sourceSum, value) => sourceSum + value, 0)
),
0,
);
const {
hidden: _hiddenSummary,
routes: _routeSummary,
...depthHeadline
} = depth.summary;
const compact = {
schemaVersion: 1,
artifacts: Object.fromEntries(
Object.entries(paths)
.filter(([name]) => name !== "output")
.map(([name, path]) => [name, artifact(path)]),
),
model: {
repo: completion.model.repo,
revision: completion.model.revision,
checkpointIdentity: completion.model.checkpoint_identity,
architecture: completion.model.architecture,
dtype: completion.model.dtype,
},
contract: {
conditions,
conditionFactors: completion.generation_contract.conditions,
edgeOrder,
edgePairs: depth.trace_contract.edge_pairs,
edgeLabels,
domains,
targetScope: depth.source_contract.target_scope,
decode: completion.generation_contract.decode,
maxNewTokens: completion.generation_contract.max_new_tokens,
useCacheGeneration: completion.generation_contract.use_cache,
useCacheTrace: depth.trace_contract.use_cache,
},
completion: {
budgetLadder: [
{
maxNewTokens: baseline.generation_contract.max_new_tokens,
naturalEos: baselineNaturalEos,
truncated: 128 - baselineNaturalEos,
outputs: 128,
},
{
maxNewTokens: 512,
naturalEos: completionEval.summary.natural_eos,
truncated: completionEval.summary.budget_truncated,
outputs: completionEval.summary.outputs,
},
],
summary: completionEval.summary,
incompleteRows,
taskRows: evaluationRows.filter(
(row) => ["code", "math"].includes(row.domain),
),
taskTotals: {
math: {
strictCorrect: mathRows.filter(
(row) => (
row.taskEvaluation.strict_complete_numeric_exact
),
).length,
total: mathRows.length,
extractionMethods: Object.fromEntries(
Object.entries(mathExtraction).map(([key, rows]) => [
key,
rows.length,
]),
),
},
code: {
testsPassed: codeRows.filter(
(row) => row.taskEvaluation.fixed_budget_tests_pass,
).length,
total: codeRows.length,
statuses: Object.fromEntries(
Object.entries(codeStatuses).map(([key, rows]) => [
key,
rows.length,
]),
),
},
},
sandbox: completionEval.sandbox,
reproduction: completionReproduction,
execution: completion.execution,
claimBoundary: completionEval.claim_boundary,
},
depth: {
summary: depthHeadline,
hiddenSeries,
routerSeries,
hiddenStages,
routerLayers,
boundaryCrossingTokens,
sources: depth.sources.map((source) => ({
id: source.id,
domain: source.domain,
label: source.label,
contentTokens: source.content_tokens,
contentTokenIdsSha256: source.content_token_ids_sha256,
boundaryCrossingTokensByCondition: (
source.boundary_crossing_tokens_by_condition
),
})),
reproduction: depthReproduction,
execution: depth.execution,
claimBoundary: depth.claim_boundary,
},
};
writeFileSync(
paths.output,
`${JSON.stringify(compact, null, 2)}\n`,
"utf8",
);
process.stdout.write(
`${paths.output}\n`
+ `${statSync(paths.output).size} bytes compact\n`
+ `${completionReproduction.generatedTokenIdsExact}`
+ ` / ${completionReproduction.cells} completion cells exact\n`
+ `${depthReproduction.sourceObjectsExactAfterRuntimeStrip}`
+ ` / ${depthReproduction.sources} depth sources exact`
+ " after runtime strip\n",
);