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.
+
+ The preview uses a 1.70 GB packed INT8 DiT instead of the
+ 3.02 GB production DiT (43.7% smaller). Switching models downloads
+ a separate DiT package. Use the same seed to compare quality.
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 dacfe27..17b1517 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 | 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
new file mode 100644
index 0000000..f371db7
--- /dev/null
+++ b/packages/acestep/optimization/experiments/OPT-0091-packed-int8-dit-runtime.md
@@ -0,0 +1,83 @@
+# 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
+
+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,