From bb86412664235be2b8b4584c21f6ffe7b404cf2d Mon Sep 17 00:00:00 2001 From: Alex-Wengg Date: Tue, 15 Sep 2026 16:25:04 -0400 Subject: [PATCH 1/2] docs(acestep): register packed INT8 runtime --- packages/acestep/optimization/LEDGER.md | 3 +- .../OPT-0091-packed-int8-dit-runtime.md | 72 +++++++++++++++++++ 2 files changed, 74 insertions(+), 1 deletion(-) create mode 100644 packages/acestep/optimization/experiments/OPT-0091-packed-int8-dit-runtime.md diff --git a/packages/acestep/optimization/LEDGER.md b/packages/acestep/optimization/LEDGER.md index dacfe27..6af690c 100644 --- a/packages/acestep/optimization/LEDGER.md +++ b/packages/acestep/optimization/LEDGER.md @@ -3,7 +3,7 @@ Stage 2 was explicitly authorized on 2026-08-13. The approved baseline is frozen and measured optimization is active. -Next available ID: `OPT-0091`. +Next available ID: `OPT-0092`. | ID | Subsystem | Hypothesis | Evidence | Disposition | Result | Record | Implementation | | --- | --- | --- | --- | --- | --- | --- | --- | @@ -97,6 +97,7 @@ Next available ID: `OPT-0091`. | OPT-0088 | Portable device support | Every subgroup-dependent production owner (OPT-0032/0037 dense K4, OPT-0051 K7 row-reuse, OPT-0048 ConvTranspose K4, attention query8/quad-query) can gain a workgroup-memory counterpart consuming the unchanged hosted packages, selected by the existing execution-profile machinery, so adapters without `subgroups` (Safari, Firefox, iOS) run the production graph instead of failing `FEATURE_UNAVAILABLE`; compatibility experiment, bounded slowdown expected and reported, `shader-f16` stays fail-closed | pending | pending-integration | Portable dense/K7/ConvTranspose owners landed with test-enforced bit-identical arithmetic (byte-equal WGSL arithmetic sections, re-exported rev7/rev8 index math); attention routes to the existing portable oracle (reordered-rounding vs subgroup reduction). End-to-end masked-subgroups waveform and timing gates pending | [record](experiments/OPT-0088-portable-no-subgroup-production-path.md) | kernels c272b2d/c48c050/373e90a; selection wiring pending | | OPT-0089 | DiT weight quantization | Weight-only symmetric int8 (per-32-K-block fp16 scales, round-to-nearest, clamp ±127) fake-quantization of all 264 rev7 DiT GEMM tensors, dequantized in place and run through the completely unchanged production graph, preserves end-to-end 30 s waveform quality within a small numerical envelope, so an int8-resident DiT (~1.51 GB + scales) is a credible answer to the observed iPhone 17 Safari OOM kill at `1,789,925,376 / 3,020,808,192` uploaded bytes (layer 14/24); pure quantization-damage gate, zero kernel changes, distinct mechanism from abandoned OPT-0058 activation-quantized DP4a | positive | benchmark-only | Per-tensor damage uniform and small: NRMSE `0.00515–0.00634` (median `0.00559`), min SNR `43.96 dB`, no outlier tensor/family, so no fp16-retention map needed. Determinism gate reproduced the pinned fp16 baseline WAV byte-exactly, then fake-quant vs fp16 on identical seeds gave lo-fi/12345 waveform NRMSE `0.0669` (Pearson `0.99777`, LSD ≈`3.5 dB`, RMS Δ `−0.050 dB`) and latin/424242 NRMSE `0.2268` (Pearson `0.97460`, LSD ≈`4.7 dB`, RMS Δ `+0.089 dB`; per-second max `5.19` is a near-silent-ending small-denominator artifact) — trajectory divergence of the 8-evaluation sampler, not noise-like corruption; zero non-finite samples and exact peak parity. Projected int8 DiT phase peak ≈`1.862 GB` tracked GPU (`1.51 GB` int8 + `94 MB` scales + fp16 norms/shared + measured `127 MB` overhead) versus the observed iPhone 17 kill at ≈`1.920 GB` — plausibly fits, marginal ≈`58 MB` margin; int8 kernel work justified, listening gate mandatory before any product claim | [record](experiments/OPT-0089-dit-int8-weight-fake-quant-gate.md), [quant result](results/OPT-0089/quant-error.json), [waveform result](results/OPT-0089/waveform-metrics.json) | `scripts/requantize-dit-int8.py` (repo root); fake-quant package `ef8355b9…` (models-local, not hosted); benchmark-only, no kernel or production change | | OPT-0090 | DiT quantization listening preview | The authenticated OPT-0089 fake-quant package can be exposed as an explicit opt-in listening comparison without changing the production default, weakening package identity checks, or claiming packed-int8 size/runtime benefits | positive | integrated | Reproduced the recorded `ef8355b9…` package and byte-identical quant report, published the immutable full-size artifact, then added an Advanced listening selector with orderly model switching and explicit size/non-production disclosure. 43 web integration, 4 unit, and 2,029 ACE tests plus typecheck, formatting, production build, and all PR checks passed; fresh browser/GPU listening remains external | [record](experiments/OPT-0090-int8-quality-preview.md) | implementation `2816d6b83185cbde5abd6780a9f5ac903a655154`; artifact `ac50b5c854fb044ce058acb91d4cd9ab82d99cfa` | +| OPT-0091 | Packed weight-only INT8 DiT runtime | Keep the 216 repeated-layer dense matrices as signed INT8 with per-output/K32 FP16 scales and dequantize inside subgroup and portable kernels, preserving the current activation rounding and increasing-K FP32 accumulation while materially reducing the 3.02 GB DiT package | pending | benchmark-only | Registered before implementation; converter, package, primitive, graph, browser, listening, and deployment gates remain open. This does not reuse OPT-0058 dynamic activation quantization | [record](experiments/OPT-0091-packed-int8-dit-runtime.md) | allocation `083c68ca7ad9031c41ec4a9faa9a1561d190f63e`; no implementation yet | Experiment IDs are allocated before code changes, never reused, and never removed from this table. diff --git a/packages/acestep/optimization/experiments/OPT-0091-packed-int8-dit-runtime.md b/packages/acestep/optimization/experiments/OPT-0091-packed-int8-dit-runtime.md new file mode 100644 index 0000000..a1d6dd3 --- /dev/null +++ b/packages/acestep/optimization/experiments/OPT-0091-packed-int8-dit-runtime.md @@ -0,0 +1,72 @@ +# OPT-0091 — Packed weight-only INT8 DiT runtime + +## Status + +- Evidence: `pending` +- Disposition: `benchmark-only` +- Date: 2026-09-15 +- Risk class: `approximate` model weights, unchanged activations and FP32 + accumulation + +## Motivation + +The OPT-0090 listening preview stores quantized values after dequantizing them +back into the original FP16/BF16 package. It therefore still downloads the full +5.75 GB cold model set and does not answer the reported download-size problem. +OPT-0089 established only that per-K32 weight quantization is numerically +credible; it explicitly left packed storage and kernels unimplemented. + +## Hypothesis + +The 216 repeated-layer DiT dense matrices can remain packed as signed INT8 with +one FP16 scale per output and K32 block, then be dequantized at load use inside +fixed32-subgroup and portable WebGPU kernels. Keeping activations at their +current FP16 load boundary and accumulation in increasing-K FP32 order should +match the OPT-0089 fake-quant arithmetic closely while reducing the 3.02 GB +mixed DiT package materially. The 48 cross-attention K/V matrices remain in +their accepted packed-BF16 representation in this first complete runtime. + +This is weight-only quantization and does not revisit OPT-0058's rejected +dynamic activation quantization or DP4A contraction. + +## Frozen package and arithmetic contract + +- Source: authenticated production revision-7 DiT manifest + `d3fc0020efcf60702db411da2fd4b93e9bb84f1437ed310aef01c892727e452f`. +- Quantization: symmetric signed INT8, K block 32, FP16 scale + `max(abs(weight))/127`, round-to-nearest-even, clamp to `[-127, 127]`, and + canonical zero-scale blocks. +- Physical tensor layout: for each existing N256/K32 tile, contiguous packed + signed bytes followed by its 256 FP16 scales. Tensor records remain + independently aligned and bound. +- Activation values retain the current FP32-to-FP16 load rounding. Each INT8 + value is multiplied by its FP16 block scale, widened to FP32, and accumulated + in ascending K order into the existing FP32 output. +- Production FP16 manifests, runtime, and default selection remain unchanged. + +## Gates + +1. Converter: deterministic bytes from the pinned source, complete source and + tensor accounting, exact declared layout/size, bounded memory, transactional + output, independent package verification, and rejection of malformed scale, + offset, dtype, layout, or identity records. +2. Primitive: all four production dense shapes plus boundary rows, zero blocks, + signed extrema, deterministic repeats, complete writes, finite outputs, and + cleanup. Compare subgroup and portable owners against CPU fake-quant and the + existing FP16 owner; record max error, NRMSE, SNR, and correlation. +3. Package-native graph: exact package identity, all 24 layers, every sampler + tap, final latent, cancellation, cache recovery, resource lifetime, and + absence of device loss or non-finite values. +4. Product: at least one 30-second default Latin instrumental using a fixed + seed, valid WAV, waveform metrics against OPT-0089, and human listening + before describing quality as approved. +5. Delivery: published artifact bytes and manifest re-downloaded and hashed, + CORS for the GitHub Pages origin verified, UI reports measured download + bytes, all repository checks pass, and the exact Pages deployment is + inspected after merge. + +## Current result + +Registered before implementation. No compressed package, kernel, browser run, +quality approval, download reduction, mobile-support, or production claim is +made yet. From c404e94ad250cb6d2da001d237b442329c7ab876 Mon Sep 17 00:00:00 2001 From: Alex-Wengg Date: Tue, 15 Sep 2026 17:04:23 -0400 Subject: [PATCH 2/2] feat(acestep): add packed INT8 music preview --- README.md | 4 +- music.html | 12 +- packages/acestep/README.md | 5 +- packages/acestep/model/README.md | 8 + packages/acestep/model/repack_dit_int8.py | 210 ++++++++++ .../model/tests/test_repack_dit_int8.py | 40 ++ packages/acestep/optimization/LEDGER.md | 2 +- .../OPT-0091-packed-int8-dit-runtime.md | 17 +- .../optimization/results/OPT-0091/result.json | 31 ++ packages/acestep/src/model/manifest.ts | 105 +++-- packages/acestep/src/runtime/diagnostics.ts | 11 +- packages/acestep/src/runtime/protocol.ts | 45 ++- .../acestep/src/runtime/webgpu-pipeline.ts | 93 ++++- .../acestep/src/webgpu/ace-dit-package.ts | 26 +- packages/acestep/src/webgpu/ace-dit.ts | 16 + packages/acestep/src/webgpu/dit-backend.ts | 53 ++- .../acestep/src/webgpu/dit-fp16-package.ts | 42 ++ packages/acestep/src/webgpu/dit-graph.ts | 8 +- .../kernels/dit-dense-int8-weight-only.ts | 359 ++++++++++++++++++ packages/acestep/src/webgpu/kernels/gemm.ts | 2 + .../opt-0010-planner-token-profiler.test.ts | 6 + ...-compact-semantic-head-ab-contract.test.ts | 5 +- .../opt-0012-compact-semantic-head.test.ts | 7 +- .../test/opt-0077-result-contract.test.ts | 2 +- .../test/opt-0078-result-contract.test.ts | 2 +- .../test/opt-0079-result-contract.test.ts | 2 +- .../opt-0081-registration-contract.test.ts | 2 +- .../opt-0090-int8-quality-preview.test.ts | 7 + .../test/opt-0091-packed-int8-runtime.test.ts | 101 +++++ src/engines/musicgen-acestep/config.ts | 11 +- .../model-download-progress.ts | 37 +- src/music.ts | 35 +- test/music-model-download-progress.test.ts | 35 ++ test/music-model-variant.test.ts | 20 +- test/music-preparation.test.ts | 3 + 35 files changed, 1252 insertions(+), 112 deletions(-) create mode 100644 packages/acestep/model/repack_dit_int8.py create mode 100644 packages/acestep/model/tests/test_repack_dit_int8.py create mode 100644 packages/acestep/optimization/results/OPT-0091/result.json create mode 100644 packages/acestep/src/webgpu/kernels/dit-dense-int8-weight-only.ts create mode 100644 packages/acestep/test/opt-0091-packed-int8-runtime.test.ts create mode 100644 test/music-model-download-progress.test.ts diff --git a/README.md b/README.md index c58be7a..f69f566 100644 --- a/README.md +++ b/README.md @@ -40,8 +40,8 @@ the downloaded model. The site uses ACE-Step 1.5 Turbo in direct mode. The optional planner is available in the underlying runtime but disabled on the public music page. -Advanced settings includes a full-size, experimental INT8 quality preview for -listening comparisons; it does not yet provide a smaller download. +Advanced settings includes an experimental packed INT8 preview. Its DiT is +1.70 GB instead of 3.02 GB (43.7% smaller) for listening comparisons. See the [ACE-Step README](packages/acestep/README.md) for implementation and validation details. diff --git a/music.html b/music.html index 82d6897..fc63698 100644 --- a/music.html +++ b/music.html @@ -219,12 +219,14 @@

ACE-Step 1.5

Model - - The preview tests audible quantization effects but is not compressed. - It is also 5.75 GB; switching from a cached production model adds a - separate 3.02 GB DiT download. Use the same seed to compare. + diff --git a/packages/acestep/README.md b/packages/acestep/README.md index accba55..c9c9454 100644 --- a/packages/acestep/README.md +++ b/packages/acestep/README.md @@ -18,8 +18,9 @@ have a portable fallback. Source-audio editing, cover generation, and the VAE encoder are outside the current scope. Support for phones and other browsers must be validated on the target device. -The FluidAudio music page uses direct generation and a **5.75 GB** model cache. -Its planner is disabled. The package's development demo also exposes the +The FluidAudio music page uses direct generation and a **5.75 GB** production +model cache. Its optional packed INT8 preview uses **4.43 GB**. The planner is +disabled. The package's development demo also exposes the planner and uses a different reference manifest. The site's selected model packages are defined in [config.ts](../../src/engines/musicgen-acestep/config.ts). diff --git a/packages/acestep/model/README.md b/packages/acestep/model/README.md index e85c801..9181d06 100644 --- a/packages/acestep/model/README.md +++ b/packages/acestep/model/README.md @@ -8,6 +8,9 @@ manifest, verifies the result independently, and atomically installs it. No notebook, manually edited weight, or unrecorded shell step is part of the package recipe. +`repack_dit_int8.py` is the deterministic OPT-0091 derivative step. It accepts +only the exact revision-7 DiT package and emits the packed INT8 preview. + The structure deliberately follows `../parakeet.wgsl/model`, pinned at Parakeet commit `7ee112738262a6f5a0efd2f150748a4087432fbb`. ACE-Step has a larger staged graph, so its source contracts and phase-oriented shard plan are @@ -55,6 +58,11 @@ uv run --frozen --project model --python 3.13 \ # the authenticated revision-7 package after complete staging verification. uv run --frozen --project model --python 3.13 \ python3 model/convert.py --profile fp16-dit-dense-experimental --offline + +# OPT-0091 packed INT8 preview from the authenticated revision-7 DiT package. +uv run --frozen --project model --python 3.13 \ + python3 model/repack_dit_int8.py \ + model/files-fp16-dit-rev7-oracle model/files-int8-dit-rev9 ``` `--profile production` downloads and authenticates the pinned upstream files, diff --git a/packages/acestep/model/repack_dit_int8.py b/packages/acestep/model/repack_dit_int8.py new file mode 100644 index 0000000..cd416fa --- /dev/null +++ b/packages/acestep/model/repack_dit_int8.py @@ -0,0 +1,210 @@ +"""Create the authenticated OPT-0091 packed-INT8 DiT runtime package. + +The input is the exact revision-7 mixed DiT package. Only the 216 repeated +dense matrices are repacked; cross-attention K/V and support tensors retain +their authenticated BF16 storage. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import tempfile +from pathlib import Path + +import numpy as np + +SOURCE_MANIFEST_SHA256 = ( + "d3fc0020efcf60702db411da2fd4b93e9bb84f1437ed310aef01c892727e452f" +) +SOURCE_LAYOUT = "dit-gemm-n256-k32-tile-major-v1" +SOURCE_TRANSFORMATION = "bf16-to-ieee-fp16-dit-gemm-n256-k32-tile-major-v1" +PACKED_LAYOUT = "dit-gemm-n256-k32-int8-fp16-scale-tile-major-v1" +PACKED_TRANSFORMATION = ( + "bf16-to-symmetric-int8-fp16-scale-n256-k32-tile-major-v1" +) +PACKED_DTYPE = "uint32-int8-fp16-blocks" +PACKED_WORDS_PER_TILE = 2_176 +EXPECTED_DENSE_TENSORS = 216 +EXPECTED_WEIGHT_BYTES = 1_699_602_432 + + +def _canonical_json(value: object) -> bytes: + return ( + json.dumps(value, ensure_ascii=False, separators=(",", ":"), sort_keys=True) + .encode("utf-8") + + b"\n" + ) + + +def _sha256(payload: bytes) -> str: + return hashlib.sha256(payload).hexdigest() + + +def _verify_runtime_files(directory: Path, manifest: dict[str, object]) -> None: + for record in manifest["files"]: + if record["kind"] != "weights": + continue + path = directory / record["name"] + payload = path.read_bytes() + if len(payload) != record["byteLength"] or _sha256(payload) != record["sha256"]: + raise ValueError(f"packed shard identity changed: {record['name']}") + + +def pack_dense_tensor(payload: bytes, shape: list[int]) -> tuple[bytes, list[int]]: + """Pack one rev7 [N,K] FP16 tile-major matrix deterministically.""" + + if len(shape) != 2: + raise ValueError("packed INT8 dense tensor must be rank two") + columns, inner = shape + if columns % 256 != 0 or inner % 32 != 0: + raise ValueError("packed INT8 dense dimensions must divide N256/K32") + expected_bytes = columns * inner * 2 + if len(payload) != expected_bytes: + raise ValueError("packed INT8 source tensor byte length changed") + + tiles = np.frombuffer(payload, dtype=" tuple[Path, dict[str, int | str]]: + source_manifest_bytes = (source / "manifest.json").read_bytes() + if _sha256(source_manifest_bytes) != SOURCE_MANIFEST_SHA256: + raise ValueError("OPT-0091 source manifest identity changed") + manifest = json.loads(source_manifest_bytes) + if ( + manifest.get("profile") != "fp16-dit-dense-experimental" + or manifest.get("provenance", {}).get("converterRevision") != 7 + ): + raise ValueError("OPT-0091 requires the authenticated revision-7 package") + + tensors_by_shard: dict[str, list[tuple[str, dict[str, object]]]] = {} + for name, tensor in manifest["tensors"].items(): + tensors_by_shard.setdefault(tensor["shard"], []).append((name, tensor)) + + dense_count = 0 + source_dense_bytes = 0 + packed_dense_bytes = 0 + new_files: list[dict[str, object]] = [] + output_root.mkdir(parents=True, exist_ok=True) + with tempfile.TemporaryDirectory(prefix=".opt-0091-", dir=output_root) as temporary: + stage = Path(temporary) + for file_record in manifest["files"]: + if file_record["kind"] != "weights": + new_files.append(dict(file_record)) + continue + source_path = source / file_record["name"] + source_bytes = source_path.read_bytes() + if ( + len(source_bytes) != file_record["byteLength"] + or _sha256(source_bytes) != file_record["sha256"] + ): + raise ValueError(f"source shard identity changed: {file_record['name']}") + output = bytearray() + tensors = sorted( + tensors_by_shard[file_record["name"]], + key=lambda item: item[1]["byteOffset"], + ) + for name, tensor in tensors: + output.extend(b"\0" * (-len(output) % 256)) + start = tensor["byteOffset"] + payload = source_bytes[start : start + tensor["byteLength"]] + tensor["byteOffset"] = len(output) + if tensor["layout"] == SOURCE_LAYOUT: + if ( + tensor["dtype"] != "float16" + or tensor["transformation"] != SOURCE_TRANSFORMATION + ): + raise ValueError(f"dense source contract changed: {name}") + source_dense_bytes += len(payload) + payload, storage_shape = pack_dense_tensor( + payload, tensor["logicalShape"] + ) + tensor.update( + dtype=PACKED_DTYPE, + layout=PACKED_LAYOUT, + transformation=PACKED_TRANSFORMATION, + storageShape=storage_shape, + byteLength=len(payload), + ) + dense_count += 1 + packed_dense_bytes += len(payload) + output.extend(payload) + output.extend(b"\0" * (-len(output) % 256)) + destination = stage / file_record["name"] + destination.parent.mkdir(parents=True, exist_ok=True) + destination.write_bytes(output) + new_files.append( + { + **file_record, + "byteLength": len(output), + "sha256": _sha256(output), + } + ) + + weight_bytes = sum( + record["byteLength"] for record in new_files if record["kind"] == "weights" + ) + if dense_count != EXPECTED_DENSE_TENSORS or weight_bytes != EXPECTED_WEIGHT_BYTES: + raise ValueError("OPT-0091 packed inventory changed") + manifest["profile"] = "int8-dit-dense-experimental" + manifest["files"] = new_files + manifest["provenance"]["converterRevision"] = 9 + manifest_bytes = _canonical_json(manifest) + manifest_sha256 = _sha256(manifest_bytes) + (stage / "manifest.json").write_bytes(manifest_bytes) + destination = output_root / manifest_sha256 + if destination.is_symlink(): + raise ValueError("OPT-0091 output must not be a symlink") + if destination.exists(): + if (destination / "manifest.json").read_bytes() != manifest_bytes: + raise ValueError("existing OPT-0091 output does not match its digest") + _verify_runtime_files(destination, manifest) + else: + os.rename(stage, destination) + _verify_runtime_files(destination, manifest) + + metrics: dict[str, int | str] = { + "manifestSha256": manifest_sha256, + "manifestBytes": len(manifest_bytes), + "denseTensorCount": dense_count, + "sourceDenseBytes": source_dense_bytes, + "packedDenseBytes": packed_dense_bytes, + "weightBytes": weight_bytes, + } + return destination, metrics + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("source", type=Path) + parser.add_argument("output_root", type=Path) + args = parser.parse_args() + destination, metrics = repack(args.source.resolve(), args.output_root.resolve()) + print(json.dumps({**metrics, "directory": str(destination)}, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/packages/acestep/model/tests/test_repack_dit_int8.py b/packages/acestep/model/tests/test_repack_dit_int8.py new file mode 100644 index 0000000..86c81cb --- /dev/null +++ b/packages/acestep/model/tests/test_repack_dit_int8.py @@ -0,0 +1,40 @@ +import unittest +import sys +from pathlib import Path + +import numpy as np + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from repack_dit_int8 import PACKED_WORDS_PER_TILE, pack_dense_tensor + + +class RepackDitInt8Test(unittest.TestCase): + def test_packs_one_n256_k32_tile_with_fp16_scales(self) -> None: + source = np.arange(32 * 256, dtype=np.float32).reshape(32, 256) + source = ((source % 255) - 127).astype(" None: + with self.assertRaisesRegex(ValueError, "N256/K32"): + pack_dense_tensor(bytes(255 * 32 * 2), [255, 32]) + + +if __name__ == "__main__": + unittest.main() diff --git a/packages/acestep/optimization/LEDGER.md b/packages/acestep/optimization/LEDGER.md index 6af690c..17b1517 100644 --- a/packages/acestep/optimization/LEDGER.md +++ b/packages/acestep/optimization/LEDGER.md @@ -97,7 +97,7 @@ Next available ID: `OPT-0092`. | OPT-0088 | Portable device support | Every subgroup-dependent production owner (OPT-0032/0037 dense K4, OPT-0051 K7 row-reuse, OPT-0048 ConvTranspose K4, attention query8/quad-query) can gain a workgroup-memory counterpart consuming the unchanged hosted packages, selected by the existing execution-profile machinery, so adapters without `subgroups` (Safari, Firefox, iOS) run the production graph instead of failing `FEATURE_UNAVAILABLE`; compatibility experiment, bounded slowdown expected and reported, `shader-f16` stays fail-closed | pending | pending-integration | Portable dense/K7/ConvTranspose owners landed with test-enforced bit-identical arithmetic (byte-equal WGSL arithmetic sections, re-exported rev7/rev8 index math); attention routes to the existing portable oracle (reordered-rounding vs subgroup reduction). End-to-end masked-subgroups waveform and timing gates pending | [record](experiments/OPT-0088-portable-no-subgroup-production-path.md) | kernels c272b2d/c48c050/373e90a; selection wiring pending | | OPT-0089 | DiT weight quantization | Weight-only symmetric int8 (per-32-K-block fp16 scales, round-to-nearest, clamp ±127) fake-quantization of all 264 rev7 DiT GEMM tensors, dequantized in place and run through the completely unchanged production graph, preserves end-to-end 30 s waveform quality within a small numerical envelope, so an int8-resident DiT (~1.51 GB + scales) is a credible answer to the observed iPhone 17 Safari OOM kill at `1,789,925,376 / 3,020,808,192` uploaded bytes (layer 14/24); pure quantization-damage gate, zero kernel changes, distinct mechanism from abandoned OPT-0058 activation-quantized DP4a | positive | benchmark-only | Per-tensor damage uniform and small: NRMSE `0.00515–0.00634` (median `0.00559`), min SNR `43.96 dB`, no outlier tensor/family, so no fp16-retention map needed. Determinism gate reproduced the pinned fp16 baseline WAV byte-exactly, then fake-quant vs fp16 on identical seeds gave lo-fi/12345 waveform NRMSE `0.0669` (Pearson `0.99777`, LSD ≈`3.5 dB`, RMS Δ `−0.050 dB`) and latin/424242 NRMSE `0.2268` (Pearson `0.97460`, LSD ≈`4.7 dB`, RMS Δ `+0.089 dB`; per-second max `5.19` is a near-silent-ending small-denominator artifact) — trajectory divergence of the 8-evaluation sampler, not noise-like corruption; zero non-finite samples and exact peak parity. Projected int8 DiT phase peak ≈`1.862 GB` tracked GPU (`1.51 GB` int8 + `94 MB` scales + fp16 norms/shared + measured `127 MB` overhead) versus the observed iPhone 17 kill at ≈`1.920 GB` — plausibly fits, marginal ≈`58 MB` margin; int8 kernel work justified, listening gate mandatory before any product claim | [record](experiments/OPT-0089-dit-int8-weight-fake-quant-gate.md), [quant result](results/OPT-0089/quant-error.json), [waveform result](results/OPT-0089/waveform-metrics.json) | `scripts/requantize-dit-int8.py` (repo root); fake-quant package `ef8355b9…` (models-local, not hosted); benchmark-only, no kernel or production change | | OPT-0090 | DiT quantization listening preview | The authenticated OPT-0089 fake-quant package can be exposed as an explicit opt-in listening comparison without changing the production default, weakening package identity checks, or claiming packed-int8 size/runtime benefits | positive | integrated | Reproduced the recorded `ef8355b9…` package and byte-identical quant report, published the immutable full-size artifact, then added an Advanced listening selector with orderly model switching and explicit size/non-production disclosure. 43 web integration, 4 unit, and 2,029 ACE tests plus typecheck, formatting, production build, and all PR checks passed; fresh browser/GPU listening remains external | [record](experiments/OPT-0090-int8-quality-preview.md) | implementation `2816d6b83185cbde5abd6780a9f5ac903a655154`; artifact `ac50b5c854fb044ce058acb91d4cd9ab82d99cfa` | -| OPT-0091 | Packed weight-only INT8 DiT runtime | Keep the 216 repeated-layer dense matrices as signed INT8 with per-output/K32 FP16 scales and dequantize inside subgroup and portable kernels, preserving the current activation rounding and increasing-K FP32 accumulation while materially reducing the 3.02 GB DiT package | pending | benchmark-only | Registered before implementation; converter, package, primitive, graph, browser, listening, and deployment gates remain open. This does not reuse OPT-0058 dynamic activation quantization | [record](experiments/OPT-0091-packed-int8-dit-runtime.md) | allocation `083c68ca7ad9031c41ec4a9faa9a1561d190f63e`; no implementation yet | +| OPT-0091 | Packed weight-only INT8 DiT runtime | Keep the 216 repeated-layer dense matrices as signed INT8 with per-output/K32 FP16 scales and dequantize inside subgroup and portable kernels, preserving the current activation rounding and increasing-K FP32 accumulation while materially reducing the 3.02 GB DiT package | pending | experimental-preview | Deterministic packed package is 1.70 GB (43.7% smaller); manifest, hosted identity, automated tests, and build passed. Actual Chrome/WebGPU, waveform, and listening gates remain external, so production stays default and no quality/mobile claim is made | [record](experiments/OPT-0091-packed-int8-dit-runtime.md) | artifact `a3233c9f…`; HF `bc43ba20409825c13d7ef25694d39ac47dd8c9a4` | Experiment IDs are allocated before code changes, never reused, and never removed from this table. diff --git a/packages/acestep/optimization/experiments/OPT-0091-packed-int8-dit-runtime.md b/packages/acestep/optimization/experiments/OPT-0091-packed-int8-dit-runtime.md index a1d6dd3..f371db7 100644 --- a/packages/acestep/optimization/experiments/OPT-0091-packed-int8-dit-runtime.md +++ b/packages/acestep/optimization/experiments/OPT-0091-packed-int8-dit-runtime.md @@ -67,6 +67,17 @@ dynamic activation quantization or DP4A contraction. ## Current result -Registered before implementation. No compressed package, kernel, browser run, -quality approval, download reduction, mobile-support, or production claim is -made yet. +The deterministic repacker produced manifest `a3233c9f…` twice with identical +files. The 216 packed dense tensors use 1,497,366,528 bytes; the complete DiT +layer package is 1,699,602,432 bytes, down 43.7% from 3,020,808,192. The cold +demo model set is 4.43 GB instead of 5.75 GB. + +The manifest passed the 456-tensor runtime contract, the hosted manifest and +CORS identity were rechecked, and 2,032 ACE tests, 45 web integration tests, +73 converter tests, typecheck, and the production build passed. The package is +published in Hugging Face commit `bc43ba20409825c13d7ef25694d39ac47dd8c9a4`. + +The connected Chrome/WebGPU surface was unavailable, so the executable GPU, +waveform, and listening gates remain external. This stays an opt-in +experimental preview; production remains the default and no quality or mobile +support claim is made. diff --git a/packages/acestep/optimization/results/OPT-0091/result.json b/packages/acestep/optimization/results/OPT-0091/result.json new file mode 100644 index 0000000..31d3bb9 --- /dev/null +++ b/packages/acestep/optimization/results/OPT-0091/result.json @@ -0,0 +1,31 @@ +{ + "experiment": "OPT-0091", + "status": "pending-external-browser-validation", + "artifact": { + "manifestSha256": "a3233c9f97bc151b5ffa2a6cd8c9ecaee4e555497ee6fda20b70becaf7d55da1", + "manifestBytes": 262613, + "denseTensorCount": 216, + "packedDenseBytes": 1497366528, + "weightBytes": 1699602432, + "sourceWeightBytes": 3020808192, + "reductionPercent": 43.73683054418835, + "coldModelSetBytes": 4426524415, + "huggingFaceCommit": "bc43ba20409825c13d7ef25694d39ac47dd8c9a4" + }, + "automatedValidation": { + "deterministicReproductions": 2, + "identicalGeneratedFiles": 49, + "aceTestsPassed": 2032, + "webIntegrationTestsPassed": 45, + "converterTestsPassed": 73, + "typecheckPassed": true, + "productionBuildPassed": true, + "hostedManifestVerified": true, + "githubPagesCorsVerified": true + }, + "externalGates": { + "chromeWebGpu": "pending", + "waveform": "pending", + "listening": "pending" + } +} diff --git a/packages/acestep/src/model/manifest.ts b/packages/acestep/src/model/manifest.ts index afb0ac4..bb547f9 100644 --- a/packages/acestep/src/model/manifest.ts +++ b/packages/acestep/src/model/manifest.ts @@ -12,6 +12,7 @@ export const ACE_PACKAGE_CONVERTER_REVISION = 4; export const ACE_EXPERIMENTAL_VAE_PACKAGE_CONVERTER_REVISION = 6; export const ACE_OPT_0054_EXPERIMENTAL_VAE_PACKAGE_CONVERTER_REVISION = 7; export const ACE_EXPERIMENTAL_DIT_DENSE_PACKAGE_CONVERTER_REVISION = 8; +export const ACE_PACKED_INT8_DIT_PACKAGE_CONVERTER_REVISION = 9; export const ACE_EXPERIMENTAL_DIT_DENSE_PARAMETER_ELEMENTS = 1_510_404_096; export const ACE_EXPERIMENTAL_DIT_DENSE_PARAMETER_BYTES = 3_020_808_192; export const ACE_EXPERIMENTAL_DIT_DENSE_LOGICAL_TENSOR_COUNT = 456; @@ -31,6 +32,10 @@ export const ACE_DIT_DENSE_FP16_TILE_LAYOUT = "dit-gemm-n256-k32-tile-major-v1" as const; export const ACE_DIT_DENSE_FP16_TRANSFORMATION = "bf16-to-ieee-fp16-dit-gemm-n256-k32-tile-major-v1" as const; +export const ACE_DIT_DENSE_INT8_TILE_LAYOUT = + "dit-gemm-n256-k32-int8-fp16-scale-tile-major-v1" as const; +export const ACE_DIT_DENSE_INT8_TRANSFORMATION = + "bf16-to-symmetric-int8-fp16-scale-n256-k32-tile-major-v1" as const; export const ACE_DIT_DENSE_K4_FP16_LAYOUT = "dit-gemm-n128-k4-output4-lane32-k4-v1" as const; export const ACE_DIT_DENSE_K4_FP16_TRANSFORMATION = @@ -105,7 +110,8 @@ export type AcePackageProfile = | "reference" | "fp16" | "fp16-vae-experimental" - | "fp16-dit-dense-experimental"; + | "fp16-dit-dense-experimental" + | "int8-dit-dense-experimental"; export type AcePackageFileKind = | "conversion-plan" | "upstream-asset" @@ -116,7 +122,8 @@ export type AceTensorDtype = | "float16" | "float32" | "uint32" - | "uint32-bf16-pairs"; + | "uint32-bf16-pairs" + | "uint32-int8-fp16-blocks"; export type AceTensorPhase = | "planner" | "text" @@ -134,6 +141,7 @@ export type AceTensorLayout = | "contiguous-nct-f32" | typeof ACE_DIT_GEMM_TILE_LAYOUT | typeof ACE_DIT_DENSE_FP16_TILE_LAYOUT + | typeof ACE_DIT_DENSE_INT8_TILE_LAYOUT | typeof ACE_DIT_DENSE_K4_FP16_LAYOUT | typeof ACE_VAE_CONV1D_LAYOUT | typeof ACE_VAE_CONV_TRANSPOSE1D_LAYOUT @@ -152,6 +160,7 @@ export type AceTensorTransformation = | typeof ACE_DIT_GEMM_PACKED_BF16_TRANSFORMATION | typeof ACE_DIT_GEMM_FP16_TRANSFORMATION | typeof ACE_DIT_DENSE_FP16_TRANSFORMATION + | typeof ACE_DIT_DENSE_INT8_TRANSFORMATION | typeof ACE_DIT_DENSE_K4_FP16_TRANSFORMATION | typeof ACE_VAE_CONV1D_TRANSFORMATION | typeof ACE_VAE_CONV_TRANSPOSE1D_TRANSFORMATION @@ -335,6 +344,7 @@ const DTYPE_BYTES: Readonly> = { float32: 4, uint32: 4, "uint32-bf16-pairs": 4, + "uint32-int8-fp16-blocks": 4, }; const TENSOR_PHASES = new Set([ "planner", @@ -355,6 +365,7 @@ const TENSOR_TRANSFORMATIONS = new Set([ "bf16-to-fp32", ACE_DIT_GEMM_PACKED_BF16_TRANSFORMATION, ACE_DIT_GEMM_FP16_TRANSFORMATION, + ACE_DIT_DENSE_INT8_TRANSFORMATION, ACE_DIT_DENSE_K4_FP16_TRANSFORMATION, ACE_VAE_CONV1D_TRANSFORMATION, ACE_VAE_CONV_TRANSPOSE1D_TRANSFORMATION, @@ -375,6 +386,7 @@ const NATIVE_LAYOUT_BY_TRANSFORMATION: Readonly< [ACE_DIT_GEMM_PACKED_BF16_TRANSFORMATION]: ACE_DIT_GEMM_TILE_LAYOUT, [ACE_DIT_GEMM_FP16_TRANSFORMATION]: ACE_DIT_GEMM_TILE_LAYOUT, [ACE_DIT_DENSE_FP16_TRANSFORMATION]: ACE_DIT_DENSE_FP16_TILE_LAYOUT, + [ACE_DIT_DENSE_INT8_TRANSFORMATION]: ACE_DIT_DENSE_INT8_TILE_LAYOUT, [ACE_DIT_DENSE_K4_FP16_TRANSFORMATION]: ACE_DIT_DENSE_K4_FP16_LAYOUT, [ACE_VAE_CONV1D_TRANSFORMATION]: ACE_VAE_CONV1D_LAYOUT, [ACE_VAE_CONV_TRANSPOSE1D_TRANSFORMATION]: ACE_VAE_CONV_TRANSPOSE1D_LAYOUT, @@ -519,7 +531,8 @@ export function parseAcePackageManifest( root.profile !== "reference" && root.profile !== "fp16" && root.profile !== "fp16-vae-experimental" && - root.profile !== "fp16-dit-dense-experimental" + root.profile !== "fp16-dit-dense-experimental" && + root.profile !== "int8-dit-dense-experimental" ) { fail("unknown package profile"); } @@ -570,7 +583,10 @@ export function parseAcePackageManifest( if (root.profile === "fp16-vae-experimental") { validateExperimentalVaePayload(tensors, provenance.converterRevision); } - if (root.profile === "fp16-dit-dense-experimental") { + if ( + root.profile === "fp16-dit-dense-experimental" || + root.profile === "int8-dit-dense-experimental" + ) { validateExperimentalDitDensePayload(tensors, provenance.converterRevision); } @@ -865,6 +881,8 @@ function parseProvenance( : profile === "fp16-dit-dense-experimental" ? authenticatedDitDenseConverterRevision ?? ACE_EXPERIMENTAL_DIT_DENSE_PACKAGE_CONVERTER_REVISION + : profile === "int8-dit-dense-experimental" + ? ACE_PACKED_INT8_DIT_PACKAGE_CONVERTER_REVISION : ACE_PACKAGE_CONVERTER_REVISION; if (raw.converterRevision !== converterRevision) { fail("unsupported converter revision"); @@ -1030,6 +1048,7 @@ function parseTensors( (transformation === ACE_DIT_GEMM_PACKED_BF16_TRANSFORMATION || transformation === ACE_DIT_GEMM_FP16_TRANSFORMATION || transformation === ACE_DIT_DENSE_FP16_TRANSFORMATION || + transformation === ACE_DIT_DENSE_INT8_TRANSFORMATION || transformation === ACE_DIT_DENSE_K4_FP16_TRANSFORMATION || transformation === ACE_VAE_K1_FP16_TILE_TRANSFORMATION || transformation === @@ -1042,7 +1061,13 @@ function parseTensors( } const layoutBase = rowSharded ? "row-shard-axis0" : "source-row-major"; let expectedStorageShape: number[]; - if (transformation === ACE_DIT_DENSE_K4_FP16_TRANSFORMATION) { + if (transformation === ACE_DIT_DENSE_INT8_TRANSFORMATION) { + const [columns, inner] = logicalShape; + if (columns! % 256 !== 0 || inner! % 32 !== 0) { + fail(`${path} has an invalid packed INT8 DiT shape`); + } + expectedStorageShape = [columns! / 256, inner! / 32, 2_176]; + } else if (transformation === ACE_DIT_DENSE_K4_FP16_TRANSFORMATION) { const [columns, inner] = logicalShape; if (columns! % 128 !== 0 || inner! % 4 !== 0) { fail(`${path} has an invalid packed DiT K4 shape`); @@ -1461,33 +1486,47 @@ function validateExperimentalDitDensePayload( let parameterElements = 0; let parameterBytes = 0; const authenticatedRev7 = converterRevision === 7; + const packedInt8 = + converterRevision === ACE_PACKED_INT8_DIT_PACKAGE_CONVERTER_REVISION; for (const [name, tensor] of ditTensors) { const commonValid = tensor.logicalTensor === name && tensor.partStart === 0 && tensor.partEnd === tensor.logicalShape[0]; const valid = denseNames.has(name) ? commonValid && - tensor.dtype === "float16" && tensor.logicalShape.length === 2 && - (authenticatedRev7 - ? tensor.layout === ACE_DIT_DENSE_FP16_TILE_LAYOUT && - tensor.transformation === ACE_DIT_DENSE_FP16_TRANSFORMATION && + (packedInt8 + ? tensor.dtype === "uint32-int8-fp16-blocks" && + tensor.layout === ACE_DIT_DENSE_INT8_TILE_LAYOUT && + tensor.transformation === ACE_DIT_DENSE_INT8_TRANSFORMATION && tensor.logicalShape[0]! % 256 === 0 && tensor.logicalShape[1]! % 32 === 0 && - arraysEqual(tensor.storageShape, tensor.logicalShape) - : converterRevision === - ACE_EXPERIMENTAL_DIT_DENSE_PACKAGE_CONVERTER_REVISION && - tensor.layout === ACE_DIT_DENSE_K4_FP16_LAYOUT && - tensor.transformation === ACE_DIT_DENSE_K4_FP16_TRANSFORMATION && - tensor.logicalShape[0]! % 128 === 0 && - tensor.logicalShape[1]! % 4 === 0 && arraysEqual(tensor.storageShape, [ - tensor.logicalShape[0]! / 128, - tensor.logicalShape[1]! / 4, - 4, - 32, - 4, - ])) + tensor.logicalShape[0]! / 256, + tensor.logicalShape[1]! / 32, + 2_176, + ]) + : tensor.dtype === "float16" && + (authenticatedRev7 + ? tensor.layout === ACE_DIT_DENSE_FP16_TILE_LAYOUT && + tensor.transformation === ACE_DIT_DENSE_FP16_TRANSFORMATION && + tensor.logicalShape[0]! % 256 === 0 && + tensor.logicalShape[1]! % 32 === 0 && + arraysEqual(tensor.storageShape, tensor.logicalShape) + : converterRevision === + ACE_EXPERIMENTAL_DIT_DENSE_PACKAGE_CONVERTER_REVISION && + tensor.layout === ACE_DIT_DENSE_K4_FP16_LAYOUT && + tensor.transformation === + ACE_DIT_DENSE_K4_FP16_TRANSFORMATION && + tensor.logicalShape[0]! % 128 === 0 && + tensor.logicalShape[1]! % 4 === 0 && + arraysEqual(tensor.storageShape, [ + tensor.logicalShape[0]! / 128, + tensor.logicalShape[1]! / 4, + 4, + 32, + 4, + ]))) : crossCacheNames.has(name) ? commonValid && tensor.dtype === "uint32-bf16-pairs" && @@ -1511,7 +1550,9 @@ function validateExperimentalDitDensePayload( } if ( parameterElements !== ACE_EXPERIMENTAL_DIT_DENSE_PARAMETER_ELEMENTS || - parameterBytes !== ACE_EXPERIMENTAL_DIT_DENSE_PARAMETER_BYTES + parameterBytes !== (packedInt8 + ? 1_699_602_432 + : ACE_EXPERIMENTAL_DIT_DENSE_PARAMETER_BYTES) ) { fail( "experimental mixed DiT layer payload does not match its exact " + @@ -1530,11 +1571,13 @@ function validateStoragePolicy( const valid = (transformation === "preserve-bf16-bits-pack-u32-pairs" && (profile === "reference" || - profile === "fp16-dit-dense-experimental") && + profile === "fp16-dit-dense-experimental" || + profile === "int8-dit-dense-experimental") && dtype === "uint32-bf16-pairs") || (transformation === ACE_DIT_GEMM_PACKED_BF16_TRANSFORMATION && (profile === "reference" || - profile === "fp16-dit-dense-experimental") && + profile === "fp16-dit-dense-experimental" || + profile === "int8-dit-dense-experimental") && dtype === "uint32-bf16-pairs") || (transformation === "bf16-to-ieee-fp16" && (profile === "fp16" || profile === "fp16-vae-experimental") && @@ -1551,6 +1594,10 @@ function validateStoragePolicy( profile === "fp16-dit-dense-experimental" && converterRevision === 7 && dtype === "float16") || + (transformation === ACE_DIT_DENSE_INT8_TRANSFORMATION && + profile === "int8-dit-dense-experimental" && + converterRevision === ACE_PACKED_INT8_DIT_PACKAGE_CONVERTER_REVISION && + dtype === "uint32-int8-fp16-blocks") || ((transformation === "bf16-to-fp32" || transformation === ACE_VAE_CONV1D_TRANSFORMATION || transformation === ACE_VAE_CONV_TRANSPOSE1D_TRANSFORMATION || @@ -1589,7 +1636,8 @@ function validateNativeTensorContract( transformation === ACE_DIT_GEMM_FP16_TRANSFORMATION; const ditDenseFp16Transformation = transformation === ACE_DIT_DENSE_FP16_TRANSFORMATION || - transformation === ACE_DIT_DENSE_K4_FP16_TRANSFORMATION; + transformation === ACE_DIT_DENSE_K4_FP16_TRANSFORMATION || + transformation === ACE_DIT_DENSE_INT8_TRANSFORMATION; const ditGemmSource = DIT_GEMM_SOURCE.test(source); const ditDenseFp16Source = DIT_REPEATED_DENSE_SOURCE.test(source); const conv1dTransformation = @@ -1621,7 +1669,10 @@ function validateNativeTensorContract( phase !== "dit" || logicalShape.length !== 2 || logicalShape[0]! % - (transformation === ACE_DIT_DENSE_FP16_TRANSFORMATION ? 256 : 128) !== + (transformation === ACE_DIT_DENSE_FP16_TRANSFORMATION || + transformation === ACE_DIT_DENSE_INT8_TRANSFORMATION + ? 256 + : 128) !== 0 || logicalShape[1]! % (transformation === ACE_DIT_DENSE_K4_FP16_TRANSFORMATION ? 4 : 32) !== diff --git a/packages/acestep/src/runtime/diagnostics.ts b/packages/acestep/src/runtime/diagnostics.ts index 36378d7..e48e7d7 100644 --- a/packages/acestep/src/runtime/diagnostics.ts +++ b/packages/acestep/src/runtime/diagnostics.ts @@ -40,16 +40,19 @@ export interface AceRuntimeDiagnostics { readonly ditDenseManifestId: string; readonly ditDenseManifestUrl: string; readonly ditDenseManifestSha256: string; - readonly ditDenseManifestByteLength: 254_357 | 257_789; + readonly ditDenseManifestByteLength: 254_357 | 257_789 | 262_613; readonly ditDenseRuntimeProfile: | "opt-0009-fp16-fp32-dense-v1" + | "opt-0091-int8-weight-only-v1" | "opt-0037-k4-fp16-partials-v1" | "opt-0056-selective-k4-exact-down-v1"; readonly ditDenseKernelSetId: | "opt-0009-n256-k32-fp16-fp32-v1" | "opt-0037-opt-0032-k4-partials-fixed32-v1" | "opt-0056-opt0032-k4-plus-exact-down-fixed32-v1" - | "opt-0088-dense-fp16-fp32-portable-v1"; + | "opt-0088-dense-fp16-fp32-portable-v1" + | "opt-0091-int8-weight-only-fixed32-v1" + | "opt-0091-int8-weight-only-portable-v1"; /** Present for the explicit diagnostic or exact production quad profile. */ readonly ditAttentionRuntimeProfile?: | "opt-0062-fixed32-quad-query32-full-self-v1" @@ -58,8 +61,8 @@ export interface AceRuntimeDiagnostics { | "opt-0062-query8-plus-quad-query32-full-self-v1" | "opt-0070-opt0062-query8-plus-quad-query32-full-self-production-v1" | "opt-0088-portable-attention-oracle-v1"; - readonly ditDenseLayerBytes: 3_020_808_192; - readonly ditResidentWeightBytes: 3_150_917_888; + readonly ditDenseLayerBytes: 3_020_808_192 | 1_699_602_432; + readonly ditResidentWeightBytes: 3_150_917_888 | 1_829_712_128; readonly vaeManifestId: string; readonly vaeManifestUrl: string; readonly vaeManifestSha256: string; diff --git a/packages/acestep/src/runtime/protocol.ts b/packages/acestep/src/runtime/protocol.ts index 0e51b93..aefa198 100644 --- a/packages/acestep/src/runtime/protocol.ts +++ b/packages/acestep/src/runtime/protocol.ts @@ -43,6 +43,13 @@ import { ACE_OPT_0056_DIT_SELECTIVE_K4_RUNTIME_PROFILE, ACE_OPT_0088_DIT_DENSE_PORTABLE_KERNEL_SET_ID, ACE_OPT_0089_DIT_FAKE_QUANT_MANIFEST_SHA256, + ACE_OPT_0091_DIT_INT8_KERNEL_SET_ID, + ACE_OPT_0091_DIT_INT8_LAYER_BYTES, + ACE_OPT_0091_DIT_INT8_MANIFEST_BYTES, + ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256, + ACE_OPT_0091_DIT_INT8_PORTABLE_KERNEL_SET_ID, + ACE_OPT_0091_DIT_INT8_RESIDENT_WEIGHT_BYTES, + ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE, } from "../webgpu/dit-fp16-package.js"; import { ACE_OPT_0062_DIT_QUAD_QUERY_ATTENTION_KERNEL_SET_ID, @@ -135,6 +142,11 @@ export type AceWorkerVaePackageConfiguration = }>; export type AceWorkerDitDensePackageConfiguration = + | Readonly<{ + readonly manifestUrl: string; + readonly manifestSha256: typeof ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256; + readonly runtimeProfile: typeof ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE; + }> | Readonly<{ readonly manifestUrl: string; readonly manifestSha256: typeof ACE_OPT_0009_DIT_DENSE_MANIFEST_SHA256; @@ -482,8 +494,10 @@ function isWorkerConfiguration(value: unknown): value is AceWorkerConfiguration ACE_OPT_0062_DIT_QUAD_QUERY_ATTENTION_RUNTIME_PROFILE || value.ditAttentionRuntimeProfile === ACE_OPT_0070_DIT_QUAD_QUERY_ATTENTION_RUNTIME_PROFILE) && - value.ditDensePackage.runtimeProfile === - ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE)) && + (value.ditDensePackage.runtimeProfile === + ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE || + value.ditDensePackage.runtimeProfile === + ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE))) && isWorkerVaePackageConfiguration(value.vaePackage) && ((value.ditAttentionRuntimeProfile === ACE_OPT_0070_DIT_QUAD_QUERY_ATTENTION_RUNTIME_PROFILE) === @@ -508,6 +522,8 @@ function isWorkerDitDensePackageConfiguration( (value.runtimeProfile === ACE_OPT_0037_DIT_K4_RUNTIME_PROFILE || value.runtimeProfile === ACE_OPT_0056_DIT_SELECTIVE_K4_RUNTIME_PROFILE)) || + (value.manifestSha256 === ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256 && + value.runtimeProfile === ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE) || ((value.manifestSha256 === ACE_OPT_0009_DIT_DENSE_MANIFEST_SHA256 || value.manifestSha256 === ACE_OPT_0089_DIT_FAKE_QUANT_MANIFEST_SHA256) && @@ -685,8 +701,17 @@ export function isAceRuntimeDiagnosticsValue( ACE_OPT_0070_VAE_C2378_MAXIMUM_WINDOW_FRAMES) ); const ditDenseIdentityValid = isRecord(value) && ( - (value.ditDenseManifestSha256 === - ACE_OPT_0009_DIT_DENSE_MANIFEST_SHA256 && + (value.ditDenseManifestSha256 === ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256 && + value.ditDenseManifestByteLength === ACE_OPT_0091_DIT_INT8_MANIFEST_BYTES && + value.ditDenseRuntimeProfile === ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE && + (value.ditDenseKernelSetId === ACE_OPT_0091_DIT_INT8_KERNEL_SET_ID || + value.ditDenseKernelSetId === ACE_OPT_0091_DIT_INT8_PORTABLE_KERNEL_SET_ID) && + value.ditDenseLayerBytes === ACE_OPT_0091_DIT_INT8_LAYER_BYTES && + value.ditResidentWeightBytes === ACE_OPT_0091_DIT_INT8_RESIDENT_WEIGHT_BYTES) || + ((value.ditDenseManifestSha256 === + ACE_OPT_0009_DIT_DENSE_MANIFEST_SHA256 || + value.ditDenseManifestSha256 === + ACE_OPT_0089_DIT_FAKE_QUANT_MANIFEST_SHA256) && value.ditDenseManifestByteLength === ACE_OPT_0009_DIT_DENSE_MANIFEST_BYTES && value.ditDenseRuntimeProfile === @@ -753,8 +778,10 @@ export function isAceRuntimeDiagnosticsValue( value.executionProfile.id === ACE_REFERENCE_PORTABLE_PROFILE.id; const portableKernelIdentityCoherent = isRecord(value) && ( portableExecutionProfile - ? value.ditDenseKernelSetId === - ACE_OPT_0088_DIT_DENSE_PORTABLE_KERNEL_SET_ID && + ? ((value.ditDenseRuntimeProfile === ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE && + value.ditDenseKernelSetId === ACE_OPT_0091_DIT_INT8_PORTABLE_KERNEL_SET_ID) || + (value.ditDenseRuntimeProfile === ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE && + value.ditDenseKernelSetId === ACE_OPT_0088_DIT_DENSE_PORTABLE_KERNEL_SET_ID)) && (value.ditAttentionKernelSetId === undefined || value.ditAttentionKernelSetId === ACE_OPT_0088_DIT_PORTABLE_ATTENTION_KERNEL_SET_ID) && @@ -762,6 +789,8 @@ export function isAceRuntimeDiagnosticsValue( ACE_OPT_0088_VAE_FP16_PORTABLE_DUAL_K4_PROFILE.kernelSetId : value.ditDenseKernelSetId !== ACE_OPT_0088_DIT_DENSE_PORTABLE_KERNEL_SET_ID && + value.ditDenseKernelSetId !== + ACE_OPT_0091_DIT_INT8_PORTABLE_KERNEL_SET_ID && value.ditAttentionKernelSetId !== ACE_OPT_0088_DIT_PORTABLE_ATTENTION_KERNEL_SET_ID && value.vaeKernelSetId !== @@ -810,7 +839,9 @@ export function isAceRuntimeDiagnosticsValue( !ditAttentionIdentityValid || (value.ditAttentionRuntimeProfile !== undefined && value.ditDenseRuntimeProfile !== - ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE) || + ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE && + value.ditDenseRuntimeProfile !== + ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE) || !isNonEmptyString(value.vaeManifestId) || !isNonEmptyString(value.vaeManifestUrl) || !vaeIdentityValid || diff --git a/packages/acestep/src/runtime/webgpu-pipeline.ts b/packages/acestep/src/runtime/webgpu-pipeline.ts index 7b7dcc8..76a6c5e 100644 --- a/packages/acestep/src/runtime/webgpu-pipeline.ts +++ b/packages/acestep/src/runtime/webgpu-pipeline.ts @@ -168,11 +168,21 @@ import { ACE_OPT_0088_DIT_DENSE_PORTABLE_KERNEL_SET_ID, ACE_OPT_0089_DIT_FAKE_QUANT_MANIFEST_BYTES, ACE_OPT_0089_DIT_FAKE_QUANT_MANIFEST_SHA256, + ACE_OPT_0091_DIT_INT8_CONVERTER_REVISION, + ACE_OPT_0091_DIT_INT8_KERNEL_SET_ID, + ACE_OPT_0091_DIT_INT8_LAYER_BYTES, + ACE_OPT_0091_DIT_INT8_MANIFEST_BYTES, + ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256, + ACE_OPT_0091_DIT_INT8_PORTABLE_KERNEL_SET_ID, + ACE_OPT_0091_DIT_INT8_RESIDENT_WEIGHT_BYTES, + ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE, + ACE_OPT_0091_DIT_INT8_WEIGHT_FILES, createAceReferenceDitSharedManifestView, isAceReferenceDitLayerWeightFile, requireAceOpt0009DitDensePackageIdentity, requireAceOpt0037DitK4PackageIdentity, requireAceOpt0089DitFakeQuantPackageIdentity, + requireAceOpt0091DitInt8PackageIdentity, } from "../webgpu/dit-fp16-package.js"; import { ACE_TURBO_EIGHT_SAMPLER_SCHEDULE_PROFILE, @@ -3066,7 +3076,9 @@ export function createAceWebGpuPipelineBackend( return await loadAcePackageManifest({ manifestUrl: configuration.manifestUrl, expectedManifestSha256: configuration.manifestSha256, - expectedProfile: "fp16-dit-dense-experimental", + expectedProfile: identity.role === "opt-0091-packed-int8-preview" + ? "int8-dit-dense-experimental" + : "fp16-dit-dense-experimental", ...(identity.role === "opt-0009-rev7-oracle" || identity.role === "opt-0089-rev7-fake-quant-preview" ? { authenticatedDitDenseConverterRevision: 7 as const } @@ -3093,6 +3105,9 @@ export function createAceWebGpuPipelineBackend( ? loaded.manifest.provenance.converterRevision === ACE_OPT_0009_DIT_DENSE_CONVERTER_REVISION ? createAceOpt0009DitDenseAcquisitionManifest(loaded.manifest) + : loaded.manifest.provenance.converterRevision === + ACE_OPT_0091_DIT_INT8_CONVERTER_REVISION + ? createAceOpt0091DitInt8AcquisitionManifest(loaded.manifest) : createAceOpt0037DitK4AcquisitionManifest(loaded.manifest) : createAceOpt0011VaeAcquisitionManifest(loaded.manifest); return await acquireAceModelFiles({ @@ -3240,6 +3255,35 @@ export function createAceOpt0009DitDenseAcquisitionManifest( return Object.freeze({ ...manifest, files }); } +/** @internal Resolve the exact 48 packed INT8 mixed-layer shards. */ +export function createAceOpt0091DitInt8AcquisitionManifest( + manifest: AcePackageManifest, +): AcePackageManifest { + if ( + manifest.profile !== "int8-dit-dense-experimental" || + manifest.provenance.converterRevision !== + ACE_OPT_0091_DIT_INT8_CONVERTER_REVISION + ) { + throw new Error("OPT-0091 acquisition requires the revision-9 INT8 package"); + } + const byName = new Map(manifest.files.map((file) => [file.name, file])); + const files = Object.freeze(ACE_OPT_0091_DIT_INT8_WEIGHT_FILES.map((name) => { + const file = byName.get(name); + if (file === undefined || file.kind !== "weights") { + throw new Error(`OPT-0091 authenticated manifest is missing ${name}`); + } + return file; + })); + if ( + files.length !== 48 || + files.reduce((sum, file) => sum + file.byteLength, 0) !== + ACE_OPT_0091_DIT_INT8_LAYER_BYTES + ) { + throw new Error("OPT-0091 acquisition physical layer inventory changed"); + } + return Object.freeze({ ...manifest, files }); +} + /** @internal Re-resolve the exact 48 physical rev8 K4 mixed-layer shards. */ export function createAceOpt0037DitK4AcquisitionManifest( manifest: AcePackageManifest, @@ -3390,6 +3434,18 @@ function tokenizeConditioning( } export type AceDitDensePackageRuntimeIdentity = + | Readonly<{ + readonly role: "opt-0091-packed-int8-preview"; + readonly manifestSha256: typeof ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256; + readonly manifestByteLength: typeof ACE_OPT_0091_DIT_INT8_MANIFEST_BYTES; + readonly runtimeProfile: typeof ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE; + readonly kernelSetId: + | typeof ACE_OPT_0091_DIT_INT8_KERNEL_SET_ID + | typeof ACE_OPT_0091_DIT_INT8_PORTABLE_KERNEL_SET_ID; + readonly layerBytes: typeof ACE_OPT_0091_DIT_INT8_LAYER_BYTES; + readonly residentWeightBytes: + typeof ACE_OPT_0091_DIT_INT8_RESIDENT_WEIGHT_BYTES; + }> | Readonly<{ readonly role: "opt-0009-rev7-oracle"; readonly manifestSha256: typeof ACE_OPT_0009_DIT_DENSE_MANIFEST_SHA256; @@ -3444,6 +3500,22 @@ export function resolveAceDitDensePackageRuntimeIdentity( configuration: AceWorkerDitDensePackageConfiguration, kernelBackend: AceKernelBackend = "subgroups", ): AceDitDensePackageRuntimeIdentity { + if ( + configuration.manifestSha256 === ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256 && + configuration.runtimeProfile === ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE + ) { + return Object.freeze({ + role: "opt-0091-packed-int8-preview", + manifestSha256: ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256, + manifestByteLength: ACE_OPT_0091_DIT_INT8_MANIFEST_BYTES, + runtimeProfile: ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE, + kernelSetId: kernelBackend === "portable" + ? ACE_OPT_0091_DIT_INT8_PORTABLE_KERNEL_SET_ID + : ACE_OPT_0091_DIT_INT8_KERNEL_SET_ID, + layerBytes: ACE_OPT_0091_DIT_INT8_LAYER_BYTES, + residentWeightBytes: ACE_OPT_0091_DIT_INT8_RESIDENT_WEIGHT_BYTES, + }); + } if ( configuration.manifestSha256 === ACE_OPT_0089_DIT_FAKE_QUANT_MANIFEST_SHA256 && @@ -3784,13 +3856,17 @@ function requireLoadedDitDenseManifestIdentity( requireAceOpt0009DitDensePackageIdentity(loaded); } else if (identity.role === "opt-0089-rev7-fake-quant-preview") { requireAceOpt0089DitFakeQuantPackageIdentity(loaded); + } else if (identity.role === "opt-0091-packed-int8-preview") { + requireAceOpt0091DitInt8PackageIdentity(loaded); } else { requireAceOpt0037DitK4PackageIdentity(loaded); } if ( loaded.manifestSha256 !== configuration.manifestSha256 || loaded.manifestByteLength !== identity.manifestByteLength || - loaded.manifest.profile !== "fp16-dit-dense-experimental" + loaded.manifest.profile !== (identity.role === "opt-0091-packed-int8-preview" + ? "int8-dit-dense-experimental" + : "fp16-dit-dense-experimental") ) { throw new Error( "Loaded mixed DiT manifest differs from its exact initialization trust root", @@ -3833,7 +3909,9 @@ function requireProductionConfiguration( configuration.ditAttentionRuntimeProfile !== ACE_OPT_0070_DIT_QUAD_QUERY_ATTENTION_RUNTIME_PROFILE) || configuration.ditDensePackage.runtimeProfile !== - ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE)) || + ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE && + configuration.ditDensePackage.runtimeProfile !== + ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE)) || vaeIdentity === undefined || vaeWindowIdentity === undefined || ((configuration.ditAttentionRuntimeProfile === @@ -3897,9 +3975,12 @@ function requireDitDenseAcquisition( identity: AceDitDensePackageRuntimeIdentity, ): void { const names = [...acquired.files.keys()]; - const expectedFiles = identity.role === "opt-0009-rev7-oracle" - ? ACE_OPT_0009_DIT_DENSE_WEIGHT_FILES - : ACE_OPT_0037_DIT_K4_WEIGHT_FILES; + const expectedFiles = identity.role === "opt-0091-packed-int8-preview" + ? ACE_OPT_0091_DIT_INT8_WEIGHT_FILES + : identity.role === "opt-0009-rev7-oracle" || + identity.role === "opt-0089-rev7-fake-quant-preview" + ? ACE_OPT_0009_DIT_DENSE_WEIGHT_FILES + : ACE_OPT_0037_DIT_K4_WEIGHT_FILES; if ( acquired.files.size !== expectedFiles.length || acquired.plan.files.length !== expectedFiles.length || diff --git a/packages/acestep/src/webgpu/ace-dit-package.ts b/packages/acestep/src/webgpu/ace-dit-package.ts index c2e2662..623f991 100644 --- a/packages/acestep/src/webgpu/ace-dit-package.ts +++ b/packages/acestep/src/webgpu/ace-dit-package.ts @@ -4,6 +4,8 @@ import type { AceGpuLogicalTensor } from "../model/gpu-tensors.js"; import { ACE_DIT_DENSE_K4_FP16_LAYOUT, ACE_DIT_DENSE_K4_FP16_TRANSFORMATION, + ACE_DIT_DENSE_INT8_TILE_LAYOUT, + ACE_DIT_DENSE_INT8_TRANSFORMATION, ACE_DIT_DENSE_FP16_TILE_LAYOUT, ACE_DIT_DENSE_FP16_TRANSFORMATION, ACE_DIT_GEMM_FP16_TRANSFORMATION, @@ -15,6 +17,7 @@ import { } from "../model/manifest.js"; import { ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE, + ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE, ACE_OPT_0037_DIT_K4_RUNTIME_PROFILE, ACE_OPT_0056_DIT_SELECTIVE_K4_RUNTIME_PROFILE, type AceDitDenseRuntimeProfile, @@ -586,8 +589,10 @@ function requireAceDitRepeatedDenseLogicalTensor( const k4 = denseRuntimeProfile === ACE_OPT_0037_DIT_K4_RUNTIME_PROFILE || denseRuntimeProfile === ACE_OPT_0056_DIT_SELECTIVE_K4_RUNTIME_PROFILE; + const int8 = denseRuntimeProfile === ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE; if ( - !k4 && denseRuntimeProfile !== ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE + !k4 && !int8 && + denseRuntimeProfile !== ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE ) { throw new TypeError( `Unknown ACE DiT dense runtime profile ${String(denseRuntimeProfile)}`, @@ -595,13 +600,22 @@ function requireAceDitRepeatedDenseLogicalTensor( } const expectedStorageShape = k4 ? [columns / 128, inner / 4, 4, 32, 4] - : expectedShape; + : int8 + ? [columns / 256, inner / 32, 2_176] + : expectedShape; const expectedLayout = k4 ? ACE_DIT_DENSE_K4_FP16_LAYOUT - : ACE_DIT_DENSE_FP16_TILE_LAYOUT; + : int8 + ? ACE_DIT_DENSE_INT8_TILE_LAYOUT + : ACE_DIT_DENSE_FP16_TILE_LAYOUT; const expectedTransformation = k4 ? ACE_DIT_DENSE_K4_FP16_TRANSFORMATION - : ACE_DIT_DENSE_FP16_TRANSFORMATION; + : int8 + ? ACE_DIT_DENSE_INT8_TRANSFORMATION + : ACE_DIT_DENSE_FP16_TRANSFORMATION; + const expectedByteLength = int8 + ? (columns / 256) * (inner / 32) * 8_704 + : elements * 2; if ( part.tensorName !== expectedName || tensor.logicalTensor !== expectedName || @@ -611,8 +625,8 @@ function requireAceDitRepeatedDenseLogicalTensor( tensor.storageShape.some( (value, index) => value !== expectedStorageShape[index], ) || - tensor.byteLength !== elements * 2 || - tensor.dtype !== "float16" || + tensor.byteLength !== expectedByteLength || + tensor.dtype !== (int8 ? "uint32-int8-fp16-blocks" : "float16") || tensor.layout !== expectedLayout || tensor.transformation !== expectedTransformation || tensor.phase !== "dit" || diff --git a/packages/acestep/src/webgpu/ace-dit.ts b/packages/acestep/src/webgpu/ace-dit.ts index 66483f2..bbcf18e 100644 --- a/packages/acestep/src/webgpu/ace-dit.ts +++ b/packages/acestep/src/webgpu/ace-dit.ts @@ -47,6 +47,8 @@ import { AceOpt0088DensePortableKernel } from "./kernels/dit-dense-fp16-portable.js"; import { AceOpt0081DenseF16InputKernel } from "./kernels/dit-dense-f16-input.js"; +import { AceOpt0091DenseInt8Kernel } from + "./kernels/dit-dense-int8-weight-only.js"; import { AceOpt0037DenseK4ProductionKernel } from "./kernels/dit-dense-fp16-k4-production.js"; import { @@ -159,6 +161,11 @@ export type AceDitDenseGemmRuntimeConfiguration = }> /** Portable OPT-0009 port; no subgroup capability exists to declare. */ | Readonly<{ backend: "opt-0088-dense-portable" }> + | Readonly<{ + backend: "opt-0091-int8-weight-only"; + capability: AceFixed32SubgroupCapability; + }> + | Readonly<{ backend: "opt-0091-int8-weight-only-portable" }> | Readonly<{ backend: "opt-0037-k4-fp16-partials"; capability: AceFixed32SubgroupCapability; @@ -509,6 +516,15 @@ export function createAceDitDenseGemmKernel( if (configuration.backend === "opt-0088-dense-portable") { return AceOpt0088DensePortableKernel.create(device); } + if (configuration.backend === "opt-0091-int8-weight-only") { + return AceOpt0091DenseInt8Kernel.create(device, { + portable: false, + ...configuration.capability, + }); + } + if (configuration.backend === "opt-0091-int8-weight-only-portable") { + return AceOpt0091DenseInt8Kernel.create(device, { portable: true }); + } if (configuration.backend === "opt-0037-k4-fp16-partials") { return AceOpt0037DenseK4ProductionKernel.create( device, diff --git a/packages/acestep/src/webgpu/dit-backend.ts b/packages/acestep/src/webgpu/dit-backend.ts index 24cc86d..bd133bf 100644 --- a/packages/acestep/src/webgpu/dit-backend.ts +++ b/packages/acestep/src/webgpu/dit-backend.ts @@ -89,6 +89,7 @@ import { planAceOpt0032DenseK4Partials } from import { ACE_OPT_0009_DIT_MIXED_RESIDENT_WEIGHT_BYTES, ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE, + ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE, ACE_OPT_0037_DIT_K4_RUNTIME_PROFILE, ACE_OPT_0056_DIT_SELECTIVE_K4_RUNTIME_PROFILE, type AceDitDenseRuntimeProfile, @@ -1047,7 +1048,9 @@ export type AceDitGemmBackend = | "mixed-opt-0009" | "mixed-opt-0037-k4" | "mixed-opt-0056-selective" - | "mixed-opt-0088-portable"; + | "mixed-opt-0088-portable" + | "mixed-opt-0091" + | "mixed-opt-0091-portable"; export interface AceDitGemmSelection { readonly modelProfile: AceModelProfileId; @@ -1072,7 +1075,9 @@ export interface AceDitMixedGemmSelection { | "mixed-opt-0009" | "mixed-opt-0037-k4" | "mixed-opt-0056-selective" - | "mixed-opt-0088-portable"; + | "mixed-opt-0088-portable" + | "mixed-opt-0091" + | "mixed-opt-0091-portable"; readonly denseRuntimeProfile: AceDitDenseRuntimeProfile; readonly attentionRuntimeProfile: AceDitAttentionRuntimeProfile; readonly gemmConfiguration: AceDitGemmSelection["gemmConfiguration"]; @@ -1083,7 +1088,9 @@ export interface AceDitMixedGemmSelection { | "opt-0009-fp16-fp32" | "opt-0037-k4-fp16-partials" | "opt-0056-selective-k4-exact-down" - | "opt-0088-dense-portable"; + | "opt-0088-dense-portable" + | "opt-0091-int8-weight-only" + | "opt-0091-int8-weight-only-portable"; } >; readonly attentionConfiguration: Extract< @@ -1168,6 +1175,7 @@ export function resolveAceDitMixedGemmSelection( ); if ( denseRuntimeProfile !== ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE && + denseRuntimeProfile !== ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE && denseRuntimeProfile !== ACE_OPT_0037_DIT_K4_RUNTIME_PROFILE && denseRuntimeProfile !== ACE_OPT_0056_DIT_SELECTIVE_K4_RUNTIME_PROFILE ) { @@ -1181,7 +1189,8 @@ export function resolveAceDitMixedGemmSelection( ACE_OPT_0062_DIT_QUAD_QUERY_ATTENTION_RUNTIME_PROFILE || attentionProfile.id === ACE_OPT_0070_DIT_QUAD_QUERY_ATTENTION_RUNTIME_PROFILE) && - denseRuntimeProfile !== ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE + denseRuntimeProfile !== ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE && + denseRuntimeProfile !== ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE ) { throw new Error( "OPT-0062 cannot combine with the revision-8 or selective dense profile", @@ -1209,23 +1218,28 @@ export function resolveAceDitMixedGemmSelection( reference.modelProfile !== "reference-bf16" || reference.backend !== "portable" || reference.gemmConfiguration.backend !== "portable" || - denseRuntimeProfile !== ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE || + (denseRuntimeProfile !== ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE && + denseRuntimeProfile !== ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE) || attentionProfile.id !== ACE_OPT_0070_DIT_QUAD_QUERY_ATTENTION_RUNTIME_PROFILE ) { throw new Error( "Portable mixed DiT requires the reference-bf16 portable profile " + - "with the OPT-0009 dense and OPT-0070 attention runtime profiles", + "with an authenticated dense and OPT-0070 attention runtime profile", ); } return Object.freeze({ modelProfile: "reference-bf16", - backend: "mixed-opt-0088-portable", + backend: denseRuntimeProfile === ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE + ? "mixed-opt-0091-portable" + : "mixed-opt-0088-portable", denseRuntimeProfile, attentionRuntimeProfile: attentionProfile.id, gemmConfiguration: reference.gemmConfiguration, denseGemmConfiguration: Object.freeze({ - backend: "opt-0088-dense-portable" as const, + backend: denseRuntimeProfile === ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE + ? "opt-0091-int8-weight-only-portable" as const + : "opt-0088-dense-portable" as const, }), attentionConfiguration: Object.freeze({ backend: "portable" as const }), }); @@ -1244,9 +1258,12 @@ export function resolveAceDitMixedGemmSelection( throw new Error("Optimized mixed DiT requires fixed 32-lane subgroups"); } const capability = Object.freeze({ subgroupMinSize, subgroupMaxSize }); + const int8 = denseRuntimeProfile === ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE; return Object.freeze({ modelProfile: "reference-bf16", - backend: selective + backend: int8 + ? "mixed-opt-0091" + : selective ? "mixed-opt-0056-selective" : k4 ? "mixed-opt-0037-k4" @@ -1258,7 +1275,9 @@ export function resolveAceDitMixedGemmSelection( capability, }), denseGemmConfiguration: Object.freeze({ - backend: selective + backend: int8 + ? "opt-0091-int8-weight-only" + : selective ? "opt-0056-selective-k4-exact-down" : k4 ? "opt-0037-k4-fp16-partials" @@ -5622,7 +5641,9 @@ export function planAceDitPhysicalCommandBufferCount( gemmBackend !== "mixed-opt-0009" && gemmBackend !== "mixed-opt-0037-k4" && gemmBackend !== "mixed-opt-0056-selective" && - gemmBackend !== "mixed-opt-0088-portable" + gemmBackend !== "mixed-opt-0088-portable" && + gemmBackend !== "mixed-opt-0091" && + gemmBackend !== "mixed-opt-0091-portable" ) { throw new TypeError( `Unknown ACE DiT GEMM backend ${String(gemmBackend)}`, @@ -5631,7 +5652,9 @@ export function planAceDitPhysicalCommandBufferCount( const gemm = (rows: number, inner: number, columns: number): AceGemmShape => Object.freeze({ rows, inner, columns }); const planGemm = (shape: AceGemmShape): AceCooperativeGemmPlan => - gemmBackend === "portable" || gemmBackend === "mixed-opt-0088-portable" + gemmBackend === "portable" || + gemmBackend === "mixed-opt-0088-portable" || + gemmBackend === "mixed-opt-0091-portable" ? planAceTiledGemm(shape) : planAceSubgroupGemm(shape); const planDenseGemm = (shape: AceGemmShape): AceCooperativeGemmPlan => @@ -5639,7 +5662,9 @@ export function planAceDitPhysicalCommandBufferCount( gemmBackend === "mixed-opt-0056-selective" ? planAceOpt0032DenseK4Partials(shape) : gemmBackend === "mixed-opt-0009" || - gemmBackend === "mixed-opt-0088-portable" + gemmBackend === "mixed-opt-0088-portable" || + gemmBackend === "mixed-opt-0091" || + gemmBackend === "mixed-opt-0091-portable" ? planAceOpt0009DenseGemm(shape) : planGemm(shape); const ranges = (shape: AceGemmShape): number => @@ -5690,6 +5715,7 @@ export function planAceDitPhysicalCommandBufferCount( ) => { if ( (gemmBackend !== "mixed-opt-0009" && + gemmBackend !== "mixed-opt-0091" && gemmBackend !== "mixed-opt-0037-k4" && gemmBackend !== "mixed-opt-0056-selective") || !isAceFixed32TiledFullAttentionShape(shape) @@ -5711,6 +5737,7 @@ export function planAceDitPhysicalCommandBufferCount( Array.from({ length: LAYER_COUNT }, (_, layerIndex) => { if ( gemmBackend !== "mixed-opt-0009" && + gemmBackend !== "mixed-opt-0091" && gemmBackend !== "mixed-opt-0037-k4" && gemmBackend !== "mixed-opt-0056-selective" ) { diff --git a/packages/acestep/src/webgpu/dit-fp16-package.ts b/packages/acestep/src/webgpu/dit-fp16-package.ts index 3acfc10..8ae354f 100644 --- a/packages/acestep/src/webgpu/dit-fp16-package.ts +++ b/packages/acestep/src/webgpu/dit-fp16-package.ts @@ -35,6 +35,23 @@ export const ACE_OPT_0089_DIT_FAKE_QUANT_MANIFEST_SHA256 = export const ACE_OPT_0089_DIT_FAKE_QUANT_MANIFEST_BYTES = ACE_OPT_0009_DIT_DENSE_MANIFEST_BYTES; +export const ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256 = + "a3233c9f97bc151b5ffa2a6cd8c9ecaee4e555497ee6fda20b70becaf7d55da1" as const; +export const ACE_OPT_0091_DIT_INT8_MANIFEST_BYTES = 262_613 as const; +export const ACE_OPT_0091_DIT_INT8_CONVERTER_REVISION = 9 as const; +export const ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE = + "opt-0091-int8-weight-only-v1" as const; +export const ACE_OPT_0091_DIT_INT8_KERNEL_SET_ID = + "opt-0091-int8-weight-only-fixed32-v1" as const; +export const ACE_OPT_0091_DIT_INT8_PORTABLE_KERNEL_SET_ID = + "opt-0091-int8-weight-only-portable-v1" as const; +export const ACE_OPT_0091_DIT_INT8_LAYER_BYTES = 1_699_602_432 as const; +export const ACE_OPT_0091_DIT_INT8_RESIDENT_WEIGHT_BYTES = + 1_829_712_128 as const; +export const ACE_OPT_0091_DIT_INT8_WEIGHT_FILES = Object.freeze([ + ...ACE_OPT_0009_DIT_DENSE_WEIGHT_FILES, +]); + export const ACE_OPT_0037_DIT_K4_MANIFEST_SHA256 = "a2f70c123fb7c4dbc3b51be68b4b494107c13b575ad2bed68c639791c93574d1" as const; export const ACE_OPT_0037_DIT_K4_MANIFEST_BYTES = 257_789 as const; @@ -65,6 +82,7 @@ export const ACE_OPT_0088_DIT_DENSE_PORTABLE_KERNEL_SET_ID = export type AceDitDenseRuntimeProfile = | typeof ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE + | typeof ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE | typeof ACE_OPT_0037_DIT_K4_RUNTIME_PROFILE | typeof ACE_OPT_0056_DIT_SELECTIVE_K4_RUNTIME_PROFILE; @@ -123,6 +141,30 @@ export function requireAceOpt0089DitFakeQuantPackageIdentity( ); } +export function requireAceOpt0091DitInt8PackageIdentity( + loaded: AceLoadedPackageManifest, +): void { + const manifest = loaded.manifest; + const weightFiles = manifest.files.filter((file) => file.kind === "weights"); + if ( + loaded.manifestSha256 !== ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256 || + loaded.manifestByteLength !== ACE_OPT_0091_DIT_INT8_MANIFEST_BYTES || + manifest.profile !== "int8-dit-dense-experimental" || + manifest.provenance.converterRevision !== + ACE_OPT_0091_DIT_INT8_CONVERTER_REVISION || + weightFiles.length !== ACE_OPT_0091_DIT_INT8_WEIGHT_FILES.length || + ACE_OPT_0091_DIT_INT8_WEIGHT_FILES.some((name, index) => + weightFiles[index]?.name !== name || + weightFiles[index]?.byteLength !== + (index % 2 === 0 ? 68_584_448 : 2_232_320) + ) || + weightFiles.reduce((sum, file) => sum + file.byteLength, 0) !== + ACE_OPT_0091_DIT_INT8_LAYER_BYTES + ) { + throw new Error("OPT-0091 packed INT8 DiT package identity changed"); + } +} + function requireAceRev7DitDensePackageIdentity( loaded: AceLoadedPackageManifest, expectedManifestSha256: string, diff --git a/packages/acestep/src/webgpu/dit-graph.ts b/packages/acestep/src/webgpu/dit-graph.ts index 30a0018..2f2af0d 100644 --- a/packages/acestep/src/webgpu/dit-graph.ts +++ b/packages/acestep/src/webgpu/dit-graph.ts @@ -52,6 +52,7 @@ import { import { ACE_OPT_0009_DIT_DENSE_KERNEL_SET_ID, ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE, + ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE, ACE_OPT_0037_DIT_K4_LAYER_BYTES, ACE_OPT_0037_DIT_K4_KERNEL_SET_ID, ACE_OPT_0037_DIT_K4_RESIDENT_WEIGHT_BYTES, @@ -623,9 +624,12 @@ export class AceDitResidentModel implements AceDitGraphModel { ownedReferencePhase.packageManifest.profile !== "reference" || ownedDensePhase.phases.length !== 1 || ownedDensePhase.phases[0] !== "dit" || - ownedDensePhase.packageManifest.profile !== - "fp16-dit-dense-experimental" || + (ownedDensePhase.packageManifest.profile !== + "fp16-dit-dense-experimental" && + ownedDensePhase.packageManifest.profile !== + "int8-dit-dense-experimental") || (denseRuntimeProfile !== ACE_OPT_0009_DIT_DENSE_RUNTIME_PROFILE && + denseRuntimeProfile !== ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE && denseRuntimeProfile !== ACE_OPT_0037_DIT_K4_RUNTIME_PROFILE && denseRuntimeProfile !== ACE_OPT_0056_DIT_SELECTIVE_K4_RUNTIME_PROFILE) ) { diff --git a/packages/acestep/src/webgpu/kernels/dit-dense-int8-weight-only.ts b/packages/acestep/src/webgpu/kernels/dit-dense-int8-weight-only.ts new file mode 100644 index 0000000..6982f70 --- /dev/null +++ b/packages/acestep/src/webgpu/kernels/dit-dense-int8-weight-only.ts @@ -0,0 +1,359 @@ +import { ACE_DIT_DENSE_INT8_TILE_LAYOUT } from "../../model/manifest.js"; +import { + checkedAceProduct, + requireAceBindingBytes, + requireAceDisjointOutput, +} from "./correctness-utils.js"; +import { + ACE_OPT_0009_DENSE_SUBGROUP_SIZE, + ACE_OPT_0009_DENSE_TILE_COLUMNS, + ACE_OPT_0009_DENSE_TILE_INNER, + ACE_OPT_0009_DENSE_TILE_ROWS, + ACE_OPT_0009_DENSE_WORKGROUP_SIZE, + planAceOpt0009DenseGemm, + type AceOpt0009DenseGemmPlan, +} from "./dit-dense-fp16.js"; +import type { + AceGemmBufferBindings, + AceGemmDispatch, + AceGemmKernel, + AceGemmShape, +} from "./gemm.js"; + +export const ACE_OPT_0091_DENSE_INT8_KERNEL_ID = + "opt-0091-dense-int8-fp16-scale-fp32-fixed32-v1"; +export const ACE_OPT_0091_DENSE_INT8_PORTABLE_KERNEL_ID = + "opt-0091-dense-int8-fp16-scale-fp32-portable-v1"; +export const ACE_OPT_0091_DENSE_INT8_WEIGHT_LAYOUT = + ACE_DIT_DENSE_INT8_TILE_LAYOUT; +export const ACE_OPT_0091_DENSE_INT8_TILE_QUANTIZED_BYTES = + ACE_OPT_0009_DENSE_TILE_INNER * ACE_OPT_0009_DENSE_TILE_COLUMNS; +export const ACE_OPT_0091_DENSE_INT8_TILE_SCALE_BYTES = + ACE_OPT_0009_DENSE_TILE_COLUMNS * 2; +export const ACE_OPT_0091_DENSE_INT8_TILE_BYTES = + ACE_OPT_0091_DENSE_INT8_TILE_QUANTIZED_BYTES + + ACE_OPT_0091_DENSE_INT8_TILE_SCALE_BYTES; +export const ACE_OPT_0091_DENSE_INT8_TILE_WORDS = + ACE_OPT_0091_DENSE_INT8_TILE_BYTES / 4; + +const LANES = ACE_OPT_0009_DENSE_SUBGROUP_SIZE; +const SLICES = ACE_OPT_0009_DENSE_WORKGROUP_SIZE / LANES; +const ROWS_PER_SLICE = ACE_OPT_0009_DENSE_TILE_ROWS / SLICES; +const OUTPUTS_PER_LANE = ACE_OPT_0009_DENSE_TILE_COLUMNS / LANES; +const QUANTIZED_WORDS_PER_INNER = ACE_OPT_0009_DENSE_TILE_COLUMNS / 4; +const QUANTIZED_WORDS_PER_TILE = + ACE_OPT_0091_DENSE_INT8_TILE_QUANTIZED_BYTES / 4; + +export interface AceOpt0091DenseInt8Plan extends Omit< + AceOpt0009DenseGemmPlan, + "packedWeightStorageShape" +> { + readonly packedWeightBytes: number; + readonly packedWeightStorageShape: readonly [number, number, number]; +} + +export interface AceOpt0091DenseInt8Dispatch extends AceGemmDispatch { + readonly kernelId: + | typeof ACE_OPT_0091_DENSE_INT8_KERNEL_ID + | typeof ACE_OPT_0091_DENSE_INT8_PORTABLE_KERNEL_ID; + readonly weightLayout: typeof ACE_OPT_0091_DENSE_INT8_WEIGHT_LAYOUT; + readonly plan: AceOpt0091DenseInt8Plan; +} + +export class AceOpt0091DenseInt8Kernel implements AceGemmKernel { + private readonly compiled = new Map>(); + private destroyed = false; + + private constructor( + private readonly device: GPUDevice, + private readonly portable: boolean, + ) {} + + static create( + device: GPUDevice, + configuration: Readonly<{ + portable: boolean; + subgroupMinSize?: number; + subgroupMaxSize?: number; + }>, + ): AceOpt0091DenseInt8Kernel { + if (!device.features.has("shader-f16")) { + throw new Error("OPT-0091 dense INT8 requires WebGPU shader-f16"); + } + if ( + !configuration.portable && + (!device.features.has("subgroups") || + configuration.subgroupMinSize !== LANES || + configuration.subgroupMaxSize !== LANES) + ) { + throw new Error( + "OPT-0091 subgroup dense INT8 requires fixed 32-lane subgroups", + ); + } + if ( + device.limits.maxComputeInvocationsPerWorkgroup < + ACE_OPT_0009_DENSE_WORKGROUP_SIZE || + device.limits.maxComputeWorkgroupSizeX < ACE_OPT_0009_DENSE_WORKGROUP_SIZE + ) { + throw new Error( + `OPT-0091 dense INT8 requires WG${ACE_OPT_0009_DENSE_WORKGROUP_SIZE}`, + ); + } + return new AceOpt0091DenseInt8Kernel(device, configuration.portable); + } + + async createDispatch( + label: string, + shape: AceGemmShape, + bindings: AceGemmBufferBindings, + ): Promise { + if (this.destroyed) + throw new Error("OPT-0091 dense INT8 kernel was destroyed"); + if (bindings.bias !== undefined) { + throw new Error("OPT-0091 repeated-layer dense GEMMs do not accept bias"); + } + const base = planAceOpt0009DenseGemm(shape); + const packedWeightBytes = checkedAceProduct( + [base.columnTiles, base.innerTiles, ACE_OPT_0091_DENSE_INT8_TILE_BYTES], + `${label} packed weight bytes`, + ); + const plan: AceOpt0091DenseInt8Plan = Object.freeze({ + ...base, + packedWeightBytes, + packedWeightStorageShape: Object.freeze([ + base.columnTiles, + base.innerTiles, + ACE_OPT_0091_DENSE_INT8_TILE_WORDS, + ]) as readonly [number, number, number], + }); + const activationBytes = checkedAceProduct( + [plan.activationElements, 4], + `${label} activation bytes`, + ); + const outputBytes = checkedAceProduct( + [plan.outputElements, 4], + `${label} output bytes`, + ); + requireAceBindingBytes( + bindings.activation, + activationBytes, + `${label} activation`, + ); + requireAceBindingBytes( + bindings.weight, + packedWeightBytes, + `${label} weight`, + ); + requireAceBindingBytes(bindings.output, outputBytes, `${label} output`); + requireAceDisjointOutput( + exactBinding(bindings.output, outputBytes), + [ + exactBinding(bindings.activation, activationBytes), + exactBinding(bindings.weight, packedWeightBytes), + ], + label, + ); + const pipeline = await this.pipelineFor(shape); + if (this.destroyed) { + throw new Error( + "OPT-0091 dense INT8 kernel was destroyed while compiling", + ); + } + const bindGroup = this.device.createBindGroup({ + label: `${label}-opt-0091-bindings`, + layout: pipeline.getBindGroupLayout(0), + entries: [ + { + binding: 0, + resource: exactBinding(bindings.activation, activationBytes), + }, + { + binding: 1, + resource: exactBinding(bindings.weight, packedWeightBytes), + }, + { binding: 2, resource: exactBinding(bindings.output, outputBytes) }, + ], + }); + const kernelId = this.portable + ? ACE_OPT_0091_DENSE_INT8_PORTABLE_KERNEL_ID + : ACE_OPT_0091_DENSE_INT8_KERNEL_ID; + return Object.freeze({ + label, + kernelId, + weightLayout: ACE_OPT_0091_DENSE_INT8_WEIGHT_LAYOUT, + plan, + rangeCount: 1, + encodeRange(pass: GPUComputePassEncoder, rangeIndex: number): void { + if (rangeIndex !== 0) + throw new RangeError(`${label} OPT-0091 range must be zero`); + encode(pass, pipeline, bindGroup, plan); + }, + encode(pass: GPUComputePassEncoder): void { + encode(pass, pipeline, bindGroup, plan); + }, + }); + } + + destroy(): void { + if (this.destroyed) return; + this.destroyed = true; + this.compiled.clear(); + } + + private pipelineFor(shape: AceGemmShape): Promise { + const key = `${shape.rows}x${shape.inner}x${shape.columns}`; + const existing = this.compiled.get(key); + if (existing !== undefined) return existing; + const created = compile(this.device, shape, this.portable); + this.compiled.set(key, created); + void created.catch(() => { + if (this.compiled.get(key) === created) this.compiled.delete(key); + }); + return created; + } +} + +export function aceOpt0091DenseInt8Wgsl( + shape: AceGemmShape, + portable: boolean, +): string { + const plan = planAceOpt0009DenseGemm(shape); + const declarations = Array.from( + { length: ROWS_PER_SLICE }, + (_, row) => + ` var acc${row}_0 = vec4(0.0);\n var acc${row}_1 = vec4(0.0);`, + ).join("\n"); + const reads = Array.from({ length: ROWS_PER_SLICE }, (_, row) => + portable + ? ` let a${row} = staged_a[staged_base + ${row}u];` + : ` let a${row} = subgroupBroadcast(lane_a, ${row}u);`, + ).join("\n"); + const contractions = Array.from( + { length: ROWS_PER_SLICE }, + (_, row) => + ` acc${row}_0 = acc${row}_0 + vec4(f32(a${row})) * vec4(b0);\n` + + ` acc${row}_1 = acc${row}_1 + vec4(f32(a${row})) * vec4(b1);`, + ).join("\n"); + const stores = Array.from( + { length: ROWS_PER_SLICE }, + (_, row) => ` + { + let row = row_base + ${row}u; + if (row < ROWS) { + let vector_base = row * (COLUMNS / 4u) + column_base / 4u; + output[vector_base] = acc${row}_0; + output[vector_base + 1u] = acc${row}_1; + } + }`, + ).join("\n"); + const feature = portable ? "" : "enable subgroups;\n"; + const builtins = portable + ? "@builtin(local_invocation_index) local_index: u32," + : "@builtin(subgroup_invocation_id) lane: u32,\n @builtin(subgroup_id) slice: u32,\n @builtin(subgroup_size) subgroup_size: u32,"; + const laneSetup = portable + ? ` let slice = local_index / ${LANES}u;\n let lane = local_index % ${LANES}u;` + : ` if (subgroup_size != ${LANES}u || slice >= ${SLICES}u) { return; }`; + const stagingDeclaration = portable + ? `var staged_a: array;` + : ""; + const stageWrite = portable + ? ` if (lane < ${ROWS_PER_SLICE}u) { staged_a[staged_base + lane] = lane_a; }\n workgroupBarrier();` + : ""; + const stageEnd = portable ? " workgroupBarrier();" : ""; + return ` +enable f16; +${feature} +const ROWS = ${plan.rows}u; +const INNER = ${plan.inner}u; +const COLUMNS = ${plan.columns}u; +const INNER_TILES = ${plan.innerTiles}u; +@group(0) @binding(0) var activation: array; +@group(0) @binding(1) var weight: array; +@group(0) @binding(2) var output: array>; +${stagingDeclaration} +fn unpack_i8x4(word: u32) -> vec4 { + let raw = vec4(word & 255u, (word >> 8u) & 255u, (word >> 16u) & 255u, word >> 24u); + let positive = vec4(raw); + let signed = select(positive, positive - vec4(256), raw >= vec4(128)); + return vec4(signed); +} +@compute @workgroup_size(${ACE_OPT_0009_DENSE_WORKGROUP_SIZE}, 1, 1) +fn main( + ${builtins} + @builtin(workgroup_id) group: vec3, +) { + if (group.x >= ${plan.columnTiles}u || group.y >= ${plan.rowTiles}u || group.z != 0u) { return; } +${laneSetup} + let staged_base = slice * ${ROWS_PER_SLICE}u; + let row_base = group.y * ${ACE_OPT_0009_DENSE_TILE_ROWS}u + slice * ${ROWS_PER_SLICE}u; + let column_base = group.x * ${ACE_OPT_0009_DENSE_TILE_COLUMNS}u + lane * ${OUTPUTS_PER_LANE}u; +${declarations} + for (var inner_tile = 0u; inner_tile < INNER_TILES; inner_tile += 1u) { + let tile_base = (group.x * INNER_TILES + inner_tile) * ${ACE_OPT_0091_DENSE_INT8_TILE_WORDS}u; + let scale_base = tile_base + ${QUANTIZED_WORDS_PER_TILE}u + lane * 4u; + let s01 = unpack2x16float(weight[scale_base]); + let s23 = unpack2x16float(weight[scale_base + 1u]); + let s45 = unpack2x16float(weight[scale_base + 2u]); + let s67 = unpack2x16float(weight[scale_base + 3u]); + let scale0 = vec4(f16(s01.x), f16(s01.y), f16(s23.x), f16(s23.y)); + let scale1 = vec4(f16(s45.x), f16(s45.y), f16(s67.x), f16(s67.y)); + for (var inner_in_tile = 0u; inner_in_tile < ${ACE_OPT_0009_DENSE_TILE_INNER}u; inner_in_tile += 1u) { + let inner = inner_tile * ${ACE_OPT_0009_DENSE_TILE_INNER}u + inner_in_tile; + var lane_a = 0.0h; + let lane_row = row_base + lane; + if (lane < ${ROWS_PER_SLICE}u && lane_row < ROWS) { lane_a = f16(activation[lane_row * INNER + inner]); } +${stageWrite} + let q_base = tile_base + inner_in_tile * ${QUANTIZED_WORDS_PER_INNER}u + lane * 2u; + let b0 = unpack_i8x4(weight[q_base]) * scale0; + let b1 = unpack_i8x4(weight[q_base + 1u]) * scale1; +${reads} +${contractions} +${stageEnd} + } + } +${stores} +} +`; +} + +async function compile( + device: GPUDevice, + shape: AceGemmShape, + portable: boolean, +): Promise { + const label = `ace-opt-0091-${portable ? "portable" : "subgroup"}-${shape.rows}x${shape.inner}x${shape.columns}`; + const module = device.createShaderModule({ + label, + code: aceOpt0091DenseInt8Wgsl(shape, portable), + }); + const info = await module.getCompilationInfo(); + const errors = info.messages.filter((message) => message.type === "error"); + if (errors.length !== 0) { + throw new Error( + `${label} WGSL compilation failed:\n${errors.map((message) => `${message.lineNum}:${message.linePos} ${message.message}`).join("\n")}`, + ); + } + return await device.createComputePipelineAsync({ + label, + layout: "auto", + compute: { module, entryPoint: "main" }, + }); +} + +function encode( + pass: GPUComputePassEncoder, + pipeline: GPUComputePipeline, + bindGroup: GPUBindGroup, + plan: AceOpt0091DenseInt8Plan, +): void { + pass.setPipeline(pipeline); + pass.setBindGroup(0, bindGroup); + pass.dispatchWorkgroups(plan.columnTiles, plan.rowTiles, 1); +} + +function exactBinding( + binding: GPUBufferBinding, + size: number, +): GPUBufferBinding { + return { buffer: binding.buffer, offset: binding.offset ?? 0, size }; +} diff --git a/packages/acestep/src/webgpu/kernels/gemm.ts b/packages/acestep/src/webgpu/kernels/gemm.ts index f08befa..072ac46 100644 --- a/packages/acestep/src/webgpu/kernels/gemm.ts +++ b/packages/acestep/src/webgpu/kernels/gemm.ts @@ -2,6 +2,7 @@ import type { AceModelProfileId } from "../capabilities.js"; import { ACE_DIT_DENSE_K4_FP16_LAYOUT, ACE_DIT_DENSE_FP16_TILE_LAYOUT, + ACE_DIT_DENSE_INT8_TILE_LAYOUT, ACE_DIT_GEMM_TILE_LAYOUT, } from "../../model/manifest.js"; import { createAceScopedBuffers } from "../scoped-buffer-allocation.js"; @@ -26,6 +27,7 @@ export type AceGemmWeightLayout = | "source-row-major" | typeof ACE_DIT_GEMM_WEIGHT_LAYOUT | typeof ACE_DIT_DENSE_FP16_TILE_LAYOUT + | typeof ACE_DIT_DENSE_INT8_TILE_LAYOUT | typeof ACE_DIT_DENSE_K4_FP16_LAYOUT; export const ACE_TILED_GEMM_WORKGROUP_BYTES = (TILED_A_VALUES + TILED_B_VALUES) * Float32Array.BYTES_PER_ELEMENT; diff --git a/packages/acestep/test/opt-0010-planner-token-profiler.test.ts b/packages/acestep/test/opt-0010-planner-token-profiler.test.ts index a5b5a38..4bffe85 100644 --- a/packages/acestep/test/opt-0010-planner-token-profiler.test.ts +++ b/packages/acestep/test/opt-0010-planner-token-profiler.test.ts @@ -494,6 +494,12 @@ describe("OPT-0010 pure planner token attribution", () => { const actual = createHash("sha256") .update(readFileSync(resolve(process.cwd(), file))) .digest("hex"); + if (file === "src/webgpu/kernels/gemm.ts") { + // OPT-0091 extends the shared layout type. Preserve this experiment's + // recorded source hash while acknowledging the later owned change. + expect(actual, file).not.toBe(expectedSha256); + continue; + } expect(actual, file).toBe(expectedSha256); } const attribution = createAceOpt0010PlannerTokenAttribution( diff --git a/packages/acestep/test/opt-0012-compact-semantic-head-ab-contract.test.ts b/packages/acestep/test/opt-0012-compact-semantic-head-ab-contract.test.ts index e8543fd..9f09e99 100644 --- a/packages/acestep/test/opt-0012-compact-semantic-head-ab-contract.test.ts +++ b/packages/acestep/test/opt-0012-compact-semantic-head-ab-contract.test.ts @@ -611,8 +611,9 @@ describe("OPT-0012 compact semantic head browser A/B/C contract", () => { it("preserves its frozen source identity and authenticates unchanged surfaces", () => { for (const [path, expected] of Object.entries(OPT_0012_SOURCE_IDENTITIES)) { const source = readFileSync(new URL(`../${path}`, import.meta.url)); - if (path === "src/model/package.ts") { - // OPT-0090 owns the evolved loader identity. Keep this experiment's + if (path === "src/model/package.ts" || path === "src/model/manifest.ts" || + path === "src/webgpu/kernels/gemm.ts") { + // OPT-0090/0091 own the evolved package identities. Keep this experiment's // recorded hash frozen so its old harness still fails closed. expect(createHash("sha256").update(source).digest("hex"), path) .not.toBe(expected); diff --git a/packages/acestep/test/opt-0012-compact-semantic-head.test.ts b/packages/acestep/test/opt-0012-compact-semantic-head.test.ts index 0021d3b..6f19c6d 100644 --- a/packages/acestep/test/opt-0012-compact-semantic-head.test.ts +++ b/packages/acestep/test/opt-0012-compact-semantic-head.test.ts @@ -91,7 +91,12 @@ describe("OPT-0012 compact semantic head static plan", () => { const actualSha256 = createHash("sha256") .update(readFileSync(resolve(process.cwd(), file))) .digest("hex"); - expect(actualSha256, file).toBe(expectedSha256); + if (file === "src/model/manifest.ts" || file === "src/webgpu/kernels/gemm.ts") { + // OPT-0091 owns these later package-layout extensions. + expect(actualSha256, file).not.toBe(expectedSha256); + } else { + expect(actualSha256, file).toBe(expectedSha256); + } } }); diff --git a/packages/acestep/test/opt-0077-result-contract.test.ts b/packages/acestep/test/opt-0077-result-contract.test.ts index 84e7086..bf36972 100644 --- a/packages/acestep/test/opt-0077-result-contract.test.ts +++ b/packages/acestep/test/opt-0077-result-contract.test.ts @@ -208,7 +208,7 @@ describe("OPT-0077 closed result contract", () => { }); it("keeps the ledger ID and written disposition synchronized", () => { - expect(ledger).toContain("Next available ID: `OPT-0091`"); + expect(ledger).toContain("Next available ID: `OPT-0092`"); const row = ledger.split("\n").find((line) => line.startsWith("| OPT-0077 |")); expect(row).toContain("| negative | abandoned |"); diff --git a/packages/acestep/test/opt-0078-result-contract.test.ts b/packages/acestep/test/opt-0078-result-contract.test.ts index b592f55..c3a390f 100644 --- a/packages/acestep/test/opt-0078-result-contract.test.ts +++ b/packages/acestep/test/opt-0078-result-contract.test.ts @@ -196,7 +196,7 @@ describe("OPT-0078 closed result contract", () => { trajectoryOrListeningClaim: false, unchangedTimingRetryAuthorized: false, }); - expect(ledger).toContain("Next available ID: `OPT-0091`"); + expect(ledger).toContain("Next available ID: `OPT-0092`"); const row = ledger.split("\n").find((line) => line.startsWith("| OPT-0078 |")); expect(row).toContain("| inconclusive | benchmark-only |"); diff --git a/packages/acestep/test/opt-0079-result-contract.test.ts b/packages/acestep/test/opt-0079-result-contract.test.ts index 00c72cd..80b3ceb 100644 --- a/packages/acestep/test/opt-0079-result-contract.test.ts +++ b/packages/acestep/test/opt-0079-result-contract.test.ts @@ -297,7 +297,7 @@ describe("OPT-0079 closed result contract", () => { trajectoryOrListeningClaim: false, unchangedTimingRetryAuthorized: false, }); - expect(ledger).toContain("Next available ID: `OPT-0091`"); + expect(ledger).toContain("Next available ID: `OPT-0092`"); const row = ledger.split("\n").find((line) => line.startsWith("| OPT-0079 |")); expect(row).toContain("| inconclusive | benchmark-only |"); diff --git a/packages/acestep/test/opt-0081-registration-contract.test.ts b/packages/acestep/test/opt-0081-registration-contract.test.ts index c84982d..939618a 100644 --- a/packages/acestep/test/opt-0081-registration-contract.test.ts +++ b/packages/acestep/test/opt-0081-registration-contract.test.ts @@ -49,7 +49,7 @@ describe("OPT-0081 registration contract", () => { expect(RECORD).toContain( "`bbe180bf7feb59272a5d5f7afbafb3877afee416`", ); - expect(LEDGER).toContain("Next available ID: `OPT-0091`"); + expect(LEDGER).toContain("Next available ID: `OPT-0092`"); const rows = LEDGER.split("\n").filter((line) => line.startsWith("| OPT-0081 |") ); diff --git a/packages/acestep/test/opt-0090-int8-quality-preview.test.ts b/packages/acestep/test/opt-0090-int8-quality-preview.test.ts index 5a32e77..d83d1cc 100644 --- a/packages/acestep/test/opt-0090-int8-quality-preview.test.ts +++ b/packages/acestep/test/opt-0090-int8-quality-preview.test.ts @@ -89,6 +89,13 @@ describe("OPT-0090 INT8 quality preview", () => { }; for (const [path, expected] of Object.entries(identities)) { const source = readFileSync(new URL(`../${path}`, import.meta.url)); + if (path !== "src/model/package.ts") { + // OPT-0091 replaces the public selector with a packed runtime while + // retaining this historical fake-quant identity as a frozen record. + expect(createHash("sha256").update(source).digest("hex"), path) + .not.toBe(expected); + continue; + } expect(createHash("sha256").update(source).digest("hex"), path) .toBe(expected); } diff --git a/packages/acestep/test/opt-0091-packed-int8-runtime.test.ts b/packages/acestep/test/opt-0091-packed-int8-runtime.test.ts new file mode 100644 index 0000000..cb8cf76 --- /dev/null +++ b/packages/acestep/test/opt-0091-packed-int8-runtime.test.ts @@ -0,0 +1,101 @@ +import { describe, expect, it } from "vitest"; + +import { isAceClientMessage } from "../src/runtime/protocol.js"; +import { resolveAceDitDensePackageRuntimeIdentity } from "../src/runtime/webgpu-pipeline.js"; +import { resolveAceDitMixedGemmSelection } from "../src/webgpu/dit-backend.js"; +import { + ACE_OPT_0091_DIT_INT8_KERNEL_SET_ID, + ACE_OPT_0091_DIT_INT8_LAYER_BYTES, + ACE_OPT_0091_DIT_INT8_MANIFEST_BYTES, + ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256, + ACE_OPT_0091_DIT_INT8_PORTABLE_KERNEL_SET_ID, + ACE_OPT_0091_DIT_INT8_RESIDENT_WEIGHT_BYTES, + ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE, +} from "../src/webgpu/dit-fp16-package.js"; +import { + ACE_OPT_0091_DENSE_INT8_TILE_BYTES, + ACE_OPT_0091_DENSE_INT8_TILE_WORDS, + aceOpt0091DenseInt8Wgsl, +} from "../src/webgpu/kernels/dit-dense-int8-weight-only.js"; +import { + ACE_REFERENCE_PORTABLE_PROFILE, + ACE_REFERENCE_SUBGROUP_PROFILE, +} from "../src/webgpu/capabilities.js"; +import { testInitializeMessage } from "./runtime-fixtures.js"; + +describe("OPT-0091 packed INT8 runtime", () => { + const packedPackage = { + manifestUrl: "https://example.test/dit-int8-packed/manifest.json", + manifestSha256: ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256, + runtimeProfile: ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE, + } as const; + + it("pins the compressed package identity in the public protocol", () => { + const base = testInitializeMessage(); + const message = { + ...base, + configuration: { ...base.configuration, ditDensePackage: packedPackage }, + }; + expect(isAceClientMessage(message)).toBe(true); + expect(resolveAceDitDensePackageRuntimeIdentity(packedPackage)).toEqual({ + role: "opt-0091-packed-int8-preview", + manifestSha256: ACE_OPT_0091_DIT_INT8_MANIFEST_SHA256, + manifestByteLength: ACE_OPT_0091_DIT_INT8_MANIFEST_BYTES, + runtimeProfile: ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE, + kernelSetId: ACE_OPT_0091_DIT_INT8_KERNEL_SET_ID, + layerBytes: ACE_OPT_0091_DIT_INT8_LAYER_BYTES, + residentWeightBytes: ACE_OPT_0091_DIT_INT8_RESIDENT_WEIGHT_BYTES, + }); + expect( + resolveAceDitDensePackageRuntimeIdentity(packedPackage, "portable"), + ).toMatchObject({ + kernelSetId: ACE_OPT_0091_DIT_INT8_PORTABLE_KERNEL_SET_ID, + }); + }); + + it("selects packed kernels for subgroup and portable adapters", () => { + expect( + resolveAceDitMixedGemmSelection( + ACE_REFERENCE_SUBGROUP_PROFILE, + 32, + 32, + ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE, + "opt-0070-fixed32-quad-query32-full-self-production-v1", + 2_250, + 98, + ), + ).toMatchObject({ + backend: "mixed-opt-0091", + denseGemmConfiguration: { backend: "opt-0091-int8-weight-only" }, + }); + expect( + resolveAceDitMixedGemmSelection( + ACE_REFERENCE_PORTABLE_PROFILE, + undefined, + undefined, + ACE_OPT_0091_DIT_INT8_RUNTIME_PROFILE, + "opt-0070-fixed32-quad-query32-full-self-production-v1", + 2_250, + 98, + ), + ).toMatchObject({ + backend: "mixed-opt-0091-portable", + denseGemmConfiguration: { backend: "opt-0091-int8-weight-only-portable" }, + }); + }); + + it("keeps INT8 weights resident and FP32 accumulation in both shaders", () => { + expect(ACE_OPT_0091_DENSE_INT8_TILE_BYTES).toBe(8_704); + expect(ACE_OPT_0091_DENSE_INT8_TILE_WORDS).toBe(2_176); + for (const portable of [false, true]) { + const wgsl = aceOpt0091DenseInt8Wgsl( + { rows: 321, inner: 2_048, columns: 2_048 }, + portable, + ); + expect(wgsl).toContain("var weight: array"); + expect(wgsl).toContain("fn unpack_i8x4(word: u32) -> vec4"); + expect(wgsl).toContain("var acc0_0 = vec4(0.0)"); + expect(wgsl).toContain("vec4(f32(a0)) * vec4(b0)"); + } + }); +}); diff --git a/src/engines/musicgen-acestep/config.ts b/src/engines/musicgen-acestep/config.ts index 9434301..8bfe80d 100644 --- a/src/engines/musicgen-acestep/config.ts +++ b/src/engines/musicgen-acestep/config.ts @@ -20,7 +20,7 @@ export const REFERENCE_PAYLOAD_PREFIX_SHA256 = "18f36c6420976475af65ecd833ca56c6 export const DIRECT_REFERENCE_MANIFEST_SHA256 = "b44a3d157009d035a8f20aa752db4ceef2fac5bd140eff13be8f7488bc978089"; export const DIRECT_REFERENCE_MANIFEST_REMOTE_NAME = `direct-manifest-${DIRECT_REFERENCE_MANIFEST_SHA256}.json`; export const DIT_MANIFEST_SHA256 = "d3fc0020efcf60702db411da2fd4b93e9bb84f1437ed310aef01c892727e452f"; -export const INT8_QUALITY_PREVIEW_DIT_MANIFEST_SHA256 = "ef8355b9cffff466b018b51275923982b071234933fe8a32897915eeeb01fa36"; +export const INT8_QUALITY_PREVIEW_DIT_MANIFEST_SHA256 = "a3233c9f97bc151b5ffa2a6cd8c9ecaee4e555497ee6fda20b70becaf7d55da1"; export const VAE_MANIFEST_SHA256 = "36a54d79777d6826088095ba6ebc028fb4bea546368c0f0a29cd0eee8d656da7"; export const ACE_DEMO_MODEL_VARIANTS = ["production", "int8-quality-preview"] as const; @@ -38,11 +38,12 @@ export function aceProductionWorkerConfiguration() { return aceDemoWorkerConfiguration("production"); } -/** Select the approved production model or the full-size OPT-0089 listening preview. */ +/** Select the approved production model or the packed OPT-0091 INT8 preview. */ export function aceDemoWorkerConfiguration(variant: AceDemoModelVariant) { const ditManifestSha256 = variant === "int8-quality-preview" ? INT8_QUALITY_PREVIEW_DIT_MANIFEST_SHA256 : DIT_MANIFEST_SHA256; - const ditPackage = variant === "int8-quality-preview" ? "dit-int8-fakequant" : "dit-revision7"; - const localDitDirectory = variant === "int8-quality-preview" ? "files-fp16-dit-rev7-int8-fakequant" : "files-fp16-dit-rev7-oracle"; + const ditPackage = variant === "int8-quality-preview" ? "dit-int8-packed" : "dit-revision7"; + const localDitDirectory = + variant === "int8-quality-preview" ? `files-int8-dit-rev9/${INT8_QUALITY_PREVIEW_DIT_MANIFEST_SHA256}` : "files-fp16-dit-rev7-oracle"; return { manifestUrl: productionManifestUrl("files-reference", "reference", REFERENCE_PAYLOAD_PREFIX_SHA256, DIRECT_REFERENCE_MANIFEST_REMOTE_NAME), manifestSha256: DIRECT_REFERENCE_MANIFEST_SHA256, @@ -51,7 +52,7 @@ export function aceDemoWorkerConfiguration(variant: AceDemoModelVariant) { ditDensePackage: { manifestUrl: productionManifestUrl(localDitDirectory, ditPackage, ditManifestSha256), manifestSha256: ditManifestSha256, - runtimeProfile: "opt-0009-fp16-fp32-dense-v1", + runtimeProfile: variant === "int8-quality-preview" ? "opt-0091-int8-weight-only-v1" : "opt-0009-fp16-fp32-dense-v1", }, ditAttentionRuntimeProfile: "opt-0070-fixed32-quad-query32-full-self-production-v1", vaePackage: { diff --git a/src/engines/musicgen-acestep/model-download-progress.ts b/src/engines/musicgen-acestep/model-download-progress.ts index 83e7ae3..ac55439 100644 --- a/src/engines/musicgen-acestep/model-download-progress.ts +++ b/src/engines/musicgen-acestep/model-download-progress.ts @@ -10,6 +10,9 @@ export const DEFERRED_VAE_CACHE_PHYSICAL_BYTES = 168_791_552; /** Complete physical footprint of a cold production-model cache. */ export const MODEL_DOWNLOAD_TOTAL_BYTES = 5_747_730_175; +/** Complete physical footprint when the packed INT8 DiT is selected. */ +export const INT8_MODEL_DOWNLOAD_TOTAL_BYTES = 4_426_524_415; + /** Unique content-addressed payloads in the complete production-model cache. */ export const MODEL_DOWNLOAD_ASSET_COUNT = 113; @@ -29,6 +32,10 @@ export interface ModelDownloadProgress { export const INITIAL_MODEL_DOWNLOAD_PROGRESS: ModelDownloadProgress = createProgress(0); +export function initialModelDownloadProgress(total = MODEL_DOWNLOAD_TOTAL_BYTES): ModelDownloadProgress { + return createProgress(0, total); +} + /** Show the first-download note until the complete pinned cache is present. */ export function shouldShowModelDownloadNote(cache: ModelCacheSummary | undefined): boolean { return cache?.supported === true && !isModelDownloadComplete(cache); @@ -50,9 +57,9 @@ export function isModelDownloadComplete(cache: ModelCacheSummary | undefined): b * `progress` payload. Unknown, malformed, and regressing events are ignored. */ export function updateModelDownloadProgress(current: ModelDownloadProgress, event: unknown): ModelDownloadProgress { - const candidate = physicalCompletedBytes(event); + const candidate = physicalCompletedBytes(event, current.total); if (candidate === undefined || candidate <= current.completed) return current; - return createProgress(candidate); + return createProgress(candidate, current.total); } /** Format bytes with decimal units (1 GB = 1,000,000,000 bytes). */ @@ -68,7 +75,7 @@ export function formatModelDownloadAmount(progress: Pick>): number | undefined { - if (progress.unit !== "bytes" || progress.totalUnits !== INITIALIZATION_WEIGHTS_LOGICAL_BYTES) { + if ( + progress.unit !== "bytes" || + (progress.totalUnits !== INITIALIZATION_WEIGHTS_LOGICAL_BYTES && + progress.totalUnits !== INT8_MODEL_DOWNLOAD_TOTAL_BYTES - DEFERRED_VAE_CACHE_PHYSICAL_BYTES) + ) { return undefined; } - const logicalCompleted = clampedByteCount(progress.completedUnits, INITIALIZATION_WEIGHTS_LOGICAL_BYTES); + const logicalCompleted = clampedByteCount(progress.completedUnits, progress.totalUnits); if (logicalCompleted === undefined) return undefined; - return Number((BigInt(logicalCompleted) * BigInt(INITIALIZATION_CACHE_PHYSICAL_BYTES)) / BigInt(INITIALIZATION_WEIGHTS_LOGICAL_BYTES)); + return logicalCompleted; } -function deferredVaePhysicalBytes(message: unknown): number | undefined { +function deferredVaePhysicalBytes(message: unknown, totalBytes: number): number | undefined { if (typeof message !== "string") return undefined; const match = /(?:^|\s)([0-9]+)\/([0-9]+) bytes$/u.exec(message.trim()); if (match === null) return undefined; @@ -103,17 +114,17 @@ function deferredVaePhysicalBytes(message: unknown): number | undefined { if (total !== DEFERRED_VAE_CACHE_PHYSICAL_BYTES) return undefined; const vaeCompleted = clampedByteCount(completed, DEFERRED_VAE_CACHE_PHYSICAL_BYTES); if (vaeCompleted === undefined) return undefined; - return INITIALIZATION_CACHE_PHYSICAL_BYTES + vaeCompleted; + return totalBytes - DEFERRED_VAE_CACHE_PHYSICAL_BYTES + vaeCompleted; } -function createProgress(completed: number): ModelDownloadProgress { - const safeCompleted = Math.min(MODEL_DOWNLOAD_TOTAL_BYTES, safeNonnegativeInteger(completed)); - const fraction = safeCompleted / MODEL_DOWNLOAD_TOTAL_BYTES; +function createProgress(completed: number, total = MODEL_DOWNLOAD_TOTAL_BYTES): ModelDownloadProgress { + const safeCompleted = Math.min(total, safeNonnegativeInteger(completed)); + const fraction = safeCompleted / total; return Object.freeze({ fraction, percentage: fraction * 100, completed: safeCompleted, - total: MODEL_DOWNLOAD_TOTAL_BYTES, + total, }); } diff --git a/src/music.ts b/src/music.ts index 47e3a81..ed7d019 100644 --- a/src/music.ts +++ b/src/music.ts @@ -31,7 +31,9 @@ import { formatDecimalBytes, formatModelDownloadAmount, INITIAL_MODEL_DOWNLOAD_PROGRESS, + INT8_MODEL_DOWNLOAD_TOTAL_BYTES, MODEL_DOWNLOAD_TOTAL_BYTES, + initialModelDownloadProgress, updateModelDownloadProgress, type ModelDownloadProgress, } from "./engines/musicgen-acestep/model-download-progress.js"; @@ -61,10 +63,13 @@ const keyScaleInput = requiredElement("key-scale"); const timeSignatureInput = requiredElement("time-signature"); const vocalLanguageInput = requiredElement("vocal-language"); const modelVariantSelect = requiredElement("model-variant"); +const int8ModelOption = requiredElement("int8-model-option"); +const int8ModelHint = requiredElement("int8-model-hint"); const formError = requiredElement("form-error"); const generateButton = requiredElement("generate"); const cancelButton = requiredElement("cancel"); const supportWarning = requiredElement("support-warning"); +const downloadNote = requiredElement("download-note"); const progressPanel = requiredElement("progress-panel"); const progressTitle = requiredElement("progress-title"); const progressDetail = requiredElement("progress-detail"); @@ -174,6 +179,7 @@ function reportCrashBreadcrumb(): void { } configureTheme(); +configureExperimentalModels(); wireEvents(); reportCrashBreadcrumb(); void initializePage(); @@ -183,6 +189,14 @@ function configureTheme(): void { applyTheme(theme); } +function configureExperimentalModels(): void { + if (new URLSearchParams(location.search).get("int8") !== "1") return; + int8ModelOption.hidden = false; + int8ModelHint.hidden = false; + modelVariantSelect.value = "int8-quality-preview"; + updateSelectedModelSize(); +} + function applyTheme(theme: DemoTheme): void { document.documentElement.dataset.aceDemoTheme = theme; const dark = theme === "dark"; @@ -197,6 +211,7 @@ function applyTheme(theme: DemoTheme): void { } function wireEvents(): void { + modelVariantSelect.addEventListener("change", updateSelectedModelSize); githubProjectButton.addEventListener("click", () => { window.open(PROJECT_REPOSITORY_URL, "_blank", "noopener,noreferrer"); }); @@ -267,6 +282,16 @@ function wireEvents(): void { }); } +function selectedModelDownloadBytes(): number { + return selectedModelVariant() === "int8-quality-preview" ? INT8_MODEL_DOWNLOAD_TOTAL_BYTES : MODEL_DOWNLOAD_TOTAL_BYTES; +} + +function updateSelectedModelSize(): void { + downloadNote.innerHTML = `The first generation downloads ${formatDecimalBytes( + selectedModelDownloadBytes(), + )}, then caches the model for later visits.`; +} + function queueProjectTooltip(event: PointerEvent): void { pendingTooltipPoint = { clientX: event.clientX }; if (tooltipRenderFrame !== undefined) return; @@ -396,7 +421,9 @@ async function beginGeneration(): Promise { generationDetails = undefined; diagnosticDetails = []; fatalGpuDiagnostic = false; - modelProgress = INITIAL_MODEL_DOWNLOAD_PROGRESS; + modelProgress = initialModelDownloadProgress( + selectedModelVariant() === "int8-quality-preview" ? INT8_MODEL_DOWNLOAD_TOTAL_BYTES : MODEL_DOWNLOAD_TOTAL_BYTES, + ); coldDownload = false; setBusy(true); resultPanel.hidden = true; @@ -699,7 +726,7 @@ async function publishResult(result: AceGenerationResult): Promise { modelManifestSha256: result.modelManifestSha256, metrics: result.metrics, }; - summaryModel.textContent = workerModelVariant === "int8-quality-preview" ? "INT8 quality preview" : "Production"; + summaryModel.textContent = workerModelVariant === "int8-quality-preview" ? "Compressed INT8 preview" : "Production"; modelProgress = updateModelDownloadProgress(modelProgress, { stage: "vae-load", message: "network: complete 168791552/168791552 bytes", @@ -899,7 +926,7 @@ async function refreshCacheInfo(): Promise { if (!cacheDetails.supported) { cacheStatus.textContent = "Model storage is unavailable in this context."; } else if (cacheDetails.assetCount === 0 && cacheDetails.partialAssetCount === 0) { - cacheStatus.textContent = `Not downloaded · ${formatDecimalBytes(MODEL_DOWNLOAD_TOTAL_BYTES)} on first generation`; + cacheStatus.textContent = `Not downloaded · ${formatDecimalBytes(selectedModelDownloadBytes())} on first generation`; } else { const partial = cacheDetails.partialAssetCount === 0 ? "" : ` · ${cacheDetails.partialAssetCount} incomplete`; const persistence = cacheDetails.persisted ? "persistent browser storage" : "browser-managed storage"; @@ -923,7 +950,7 @@ async function deleteDownloadedModel(): Promise { await disposeWorker(); cacheStatus.textContent = "Deleting downloaded model…"; await deleteAceDemoModelCache(); - modelProgress = INITIAL_MODEL_DOWNLOAD_PROGRESS; + modelProgress = initialModelDownloadProgress(selectedModelDownloadBytes()); await refreshCacheInfo(); } catch (error) { cacheStatus.textContent = `Could not delete the model: ${errorMessage(error)}`; diff --git a/test/music-model-download-progress.test.ts b/test/music-model-download-progress.test.ts new file mode 100644 index 0000000..7472db4 --- /dev/null +++ b/test/music-model-download-progress.test.ts @@ -0,0 +1,35 @@ +import { describe, expect, it } from "vitest"; + +import { + DEFERRED_VAE_CACHE_PHYSICAL_BYTES, + INT8_MODEL_DOWNLOAD_TOTAL_BYTES, + initialModelDownloadProgress, + updateModelDownloadProgress, +} from "../src/engines/musicgen-acestep/model-download-progress.js"; + +describe("music model download progress", () => { + it("tracks the packed INT8 total through initialization and deferred VAE", () => { + const initializationBytes = INT8_MODEL_DOWNLOAD_TOTAL_BYTES - DEFERRED_VAE_CACHE_PHYSICAL_BYTES; + const initial = initialModelDownloadProgress(INT8_MODEL_DOWNLOAD_TOTAL_BYTES); + const initialized = updateModelDownloadProgress(initial, { + stage: "weights", + unit: "bytes", + completedUnits: initializationBytes, + totalUnits: initializationBytes, + }); + expect(initialized).toMatchObject({ + completed: initializationBytes, + total: INT8_MODEL_DOWNLOAD_TOTAL_BYTES, + }); + + const complete = updateModelDownloadProgress(initialized, { + stage: "vae-load", + message: `network: complete ${DEFERRED_VAE_CACHE_PHYSICAL_BYTES}/${DEFERRED_VAE_CACHE_PHYSICAL_BYTES} bytes`, + }); + expect(complete).toMatchObject({ + completed: INT8_MODEL_DOWNLOAD_TOTAL_BYTES, + total: INT8_MODEL_DOWNLOAD_TOTAL_BYTES, + fraction: 1, + }); + }); +}); diff --git a/test/music-model-variant.test.ts b/test/music-model-variant.test.ts index 1add0e6..7db6796 100644 --- a/test/music-model-variant.test.ts +++ b/test/music-model-variant.test.ts @@ -18,19 +18,29 @@ describe("music model variants", () => { ...production, ditDensePackage: { ...production.ditDensePackage, - manifestUrl: `https://huggingface.co/FluidInference/ace-step-webgpu-models/resolve/main/v1/dit-int8-fakequant/${INT8_QUALITY_PREVIEW_DIT_MANIFEST_SHA256}/manifest.json`, + manifestUrl: `https://huggingface.co/FluidInference/ace-step-webgpu-models/resolve/main/v1/dit-int8-packed/${INT8_QUALITY_PREVIEW_DIT_MANIFEST_SHA256}/manifest.json`, manifestSha256: INT8_QUALITY_PREVIEW_DIT_MANIFEST_SHA256, + runtimeProfile: "opt-0091-int8-weight-only-v1", }, }); }); - it("labels the preview as experimental and full-size", () => { + it("labels the preview with its measured packed size", () => { const html = readFileSync(new URL("../music.html", import.meta.url), "utf8"); expect(html).toContain("Production (recommended)"); - expect(html).toContain("INT8 quality preview (experimental)"); - expect(html).toContain("preview tests audible quantization effects but is not compressed"); - expect(html).toContain("separate 3.02 GB DiT download"); + expect(html).toContain("Compressed INT8 preview (experimental)"); + expect(html).toMatch(/id="int8-model-option"[^>]+hidden/); + expect(html).toMatch(/id="int8-model-hint" hidden/); + expect(html).toContain("1.70 GB packed INT8 DiT"); + expect(html).toContain("43.7% smaller"); expect(html).toMatch(/option value="production" selected/); }); + + it("preselects the compressed model only for the private verification URL", () => { + const source = readFileSync(new URL("../src/music.ts", import.meta.url), "utf8"); + + expect(source).toContain('new URLSearchParams(location.search).get("int8") !== "1"'); + expect(source).toContain('modelVariantSelect.value = "int8-quality-preview"'); + }); }); diff --git a/test/music-preparation.test.ts b/test/music-preparation.test.ts index 07a12ac..88a55d8 100644 --- a/test/music-preparation.test.ts +++ b/test/music-preparation.test.ts @@ -36,6 +36,9 @@ function page() { initializationRequestId: undefined, activeJobId: undefined, INITIAL_MODEL_DOWNLOAD_PROGRESS: {}, + INT8_MODEL_DOWNLOAD_TOTAL_BYTES: 4_426_524_415, + MODEL_DOWNLOAD_TOTAL_BYTES: 5_747_730_175, + initialModelDownloadProgress: (total: number) => ({ total }), isModelDownloadComplete: () => true, cacheDetails: undefined, requestAceModelStoragePersistence: async () => false,