Skip to content

Commit 9a34ac6

Browse files
committed
feat(qwen3-a8w8): wire serving backend
1 parent 8c0b8f9 commit 9a34ac6

5 files changed

Lines changed: 779 additions & 48 deletions

File tree

examples/model/qwen3_14b/npu_generate.py

Lines changed: 83 additions & 47 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
from __future__ import annotations
1111

1212
import argparse
13+
import dataclasses
1314
import statistics
1415
import sys
1516
import time
@@ -33,11 +34,17 @@ def _bootstrap_package_root() -> None:
3334

3435
from python.core import GenerateConfig, LLMEngine, RuntimeConfig
3536
from python.core.kv_cache import KvCacheManager
37+
from python.core.model_loader import ModelLoader
3638
from python.core.parallel import ParallelConfig, parse_device_ids
3739
from python.profile import get_profiler, merge_profile, profile_span
38-
from examples.model.qwen3_14b.runner.npu_executor import Qwen314BPyptoExecutor as PyptoExecutor
40+
from examples.model.qwen3_14b.runner.a8w8_loader import Qwen3A8W8DirectoryLoader
41+
from examples.model.qwen3_14b.runner.npu_executor import Qwen314BPyptoExecutor
42+
from examples.model.qwen3_14b.runner.npu_executor_a8w8 import Qwen314BA8W8PyptoExecutor
3943
from python.core.types import LoadedModel
40-
import dataclasses
44+
45+
46+
_QWEN3_BF16_FORMAT = "qwen3-14b"
47+
_QWEN3_A8W8_FORMAT = "qwen3-a8w8"
4148

4249

4350
# -----------------------------------------------------------------------------
@@ -70,24 +77,6 @@ def TimePhase(self, name: str):
7077
finally:
7178
self.phases[name] = self.phases.get(name, 0.0) + (time.perf_counter() - t0)
7279

73-
def WrapKernel(self, fn, name: str, *, group_by_decode_step: bool = False):
74-
"""Return a wrapper that records every call's duration under `name`."""
75-
76-
def wrapper(*args, **kwargs):
77-
t0 = time.perf_counter()
78-
try:
79-
return fn(*args, **kwargs)
80-
finally:
81-
dt = time.perf_counter() - t0
82-
self.kernel_times[name].append(dt)
83-
if group_by_decode_step and self._decode_step_idx >= 0:
84-
bucket = self.kernel_per_decode_step[name]
85-
while len(bucket) <= self._decode_step_idx:
86-
bucket.append([])
87-
bucket[self._decode_step_idx].append(dt)
88-
89-
return wrapper
90-
9180
def BeginDecodeStep(self) -> None:
9281
self._decode_step_idx += 1
9382

@@ -145,24 +134,23 @@ def InstallProfiling(engine: LLMEngine, model_id: str, collector: _TimingCollect
145134
"""
146135
executor = engine._executor # type: ignore[attr-defined]
147136
compiled = executor._compiled[model_id] # type: ignore[attr-defined]
137+
runner = executor._runners[model_id] # type: ignore[attr-defined]
138+
kernel_names = {
139+
id(compiled.prefill): ("kernel.prefill_fwd", False),
140+
id(compiled.decode): ("kernel.decode_layer", True),
141+
}
148142

149-
# Kernels are dispatched by Qwen314BModelRunner.
150-
if hasattr(compiled.prefill, "chip_callable") or hasattr(compiled.prefill, "compiled"):
151-
runner = executor._runners[model_id] # type: ignore[attr-defined]
152-
orig_run_program = runner._run_distributed_program # type: ignore[attr-defined]
153-
kernel_names = {
154-
id(compiled.prefill): ("kernel.prefill_fwd", False),
155-
id(compiled.decode): ("kernel.decode_layer", True),
156-
}
143+
def install_runner_kernel_timing(method_name: str) -> None:
144+
orig_run = getattr(runner, method_name)
157145

158-
def timed_run_program(callable_spec, *args, **kwargs):
146+
def timed_run(callable_spec, *args, **kwargs):
159147
kernel_info = kernel_names.get(id(callable_spec))
160148
if kernel_info is None:
161-
return orig_run_program(callable_spec, *args, **kwargs)
149+
return orig_run(callable_spec, *args, **kwargs)
162150
name, group_by_decode_step = kernel_info
163151
t0 = time.perf_counter()
164152
try:
165-
timing = orig_run_program(callable_spec, *args, **kwargs)
153+
timing = orig_run(callable_spec, *args, **kwargs)
166154
finally:
167155
dt = time.perf_counter() - t0
168156
collector.kernel_times[name].append(dt)
@@ -174,14 +162,16 @@ def timed_run_program(callable_spec, *args, **kwargs):
174162
collector.RecordRunTiming(name, timing)
175163
return timing
176164

177-
runner._run_distributed_program = timed_run_program # type: ignore[attr-defined]
165+
setattr(runner, method_name, timed_run)
166+
167+
# A8W8 kernels are dispatched as L2 callables by Qwen314BA8W8ModelRunner.
168+
if hasattr(compiled.prefill, "chip_callable"):
169+
install_runner_kernel_timing("_run_l2_program")
170+
# Original Qwen3-14B kernels are dispatched by the L3 model runner.
171+
elif hasattr(compiled.prefill, "compiled"):
172+
install_runner_kernel_timing("_run_distributed_program")
178173
else:
179-
# Per-layer kernel wrappers. compiled.prefill / compiled.decode are invoked
180-
# once per transformer layer inside run_prefill / run_decode respectively.
181-
compiled.prefill = collector.WrapKernel(compiled.prefill, "kernel.prefill_fwd")
182-
compiled.decode = collector.WrapKernel(
183-
compiled.decode, "kernel.decode_layer", group_by_decode_step=True
184-
)
174+
raise TypeError("unsupported compiled kernel wrapper for profiling")
185175

186176
# Top-level executor API wrappers.
187177
orig_prefill = executor.run_prefill
@@ -345,6 +335,15 @@ def build_parser() -> argparse.ArgumentParser:
345335
parser.add_argument("--model-dir", required=True, help="Local model directory, e.g. a Hugging Face snapshot.")
346336
parser.add_argument("--prompt", required=True, help="Prompt text.")
347337
parser.add_argument("--model-id", default="qwen3-14b-local")
338+
parser.add_argument(
339+
"--model-format",
340+
default=_QWEN3_BF16_FORMAT,
341+
choices=[_QWEN3_BF16_FORMAT, _QWEN3_A8W8_FORMAT],
342+
help=(
343+
"Qwen3-14B weight format. Use qwen3-14b for the original BF16/L3 path "
344+
"or qwen3-a8w8 for compressed-tensors W8A8 checkpoints."
345+
),
346+
)
348347
parser.add_argument("--platform", default="a2a3", choices=["a2a3sim", "a2a3", "a5sim", "a5"])
349348
parser.add_argument("--device-id", type=int, default=0, help="Default NPU device id when --devices is unset.")
350349
parser.add_argument(
@@ -367,12 +366,24 @@ def build_parser() -> argparse.ArgumentParser:
367366
help="Offline generation does not launch DP replicas; values > 1 fail fast.",
368367
)
369368
parser.add_argument("--max-seq-len", type=int, default=4096)
369+
parser.add_argument("--max-batch-size", type=int, default=16)
370370
parser.add_argument("--max-new-tokens", type=int, default=32)
371+
parser.add_argument(
372+
"--decode-backend",
373+
default="a8w8",
374+
choices=["a8w8"],
375+
help="For qwen3-a8w8 only: run the A8W8 prefill/decode backend.",
376+
)
371377
parser.add_argument("--temperature", type=float, default=0.0)
372378
parser.add_argument("--top-p", type=float, default=1.0)
373379
parser.add_argument("--top-k", type=int, default=None)
374380
parser.add_argument("--stream", action="store_true", default=False)
375381
parser.add_argument("--save-kernels-dir", default=None)
382+
parser.add_argument(
383+
"--pto-isa-commit",
384+
default=None,
385+
help="For qwen3-a8w8 only: pin PyPTO compile/assemble to the installed runtime's pto-isa revision.",
386+
)
376387
parser.add_argument(
377388
"--num-layers-override",
378389
type=int,
@@ -395,6 +406,22 @@ def build_parser() -> argparse.ArgumentParser:
395406
return parser
396407

397408

409+
def _model_loader_for_format(model_format: str) -> ModelLoader | None:
410+
if model_format != _QWEN3_A8W8_FORMAT:
411+
return None
412+
model_loader = ModelLoader()
413+
model_loader.register(Qwen3A8W8DirectoryLoader())
414+
return model_loader
415+
416+
417+
def _executor_class_for_format(model_format: str):
418+
if model_format == _QWEN3_A8W8_FORMAT:
419+
return Qwen314BA8W8PyptoExecutor
420+
if model_format == _QWEN3_BF16_FORMAT:
421+
return Qwen314BPyptoExecutor
422+
raise ValueError(f"unsupported model_format: {model_format!r}")
423+
424+
398425
def main() -> None:
399426
args = build_parser().parse_args()
400427
get_profiler(process_name="npu_generate")
@@ -417,14 +444,22 @@ def main() -> None:
417444
device_ids = parallel_config.replica_device_groups[0]
418445

419446
kv_cache_manager = KvCacheManager()
420-
executor = PyptoExecutor(
447+
executor_cls = _executor_class_for_format(args.model_format)
448+
executor_kwargs = {
449+
"platform": args.platform,
450+
"device_ids": device_ids,
451+
"save_kernels_dir": args.save_kernels_dir,
452+
"l3_trace": args.profile_verbose,
453+
}
454+
if args.model_format == _QWEN3_A8W8_FORMAT:
455+
executor_kwargs["pto_isa_commit"] = args.pto_isa_commit
456+
executor = executor_cls(
421457
kv_cache_manager,
422-
platform=args.platform,
423-
device_ids=device_ids,
424-
save_kernels_dir=args.save_kernels_dir,
425-
l3_trace=args.profile_verbose,
458+
**executor_kwargs,
426459
)
460+
model_loader = _model_loader_for_format(args.model_format)
427461
engine = LLMEngine(
462+
model_loader=model_loader,
428463
kv_cache_manager=kv_cache_manager,
429464
executor=executor,
430465
)
@@ -437,15 +472,16 @@ def main() -> None:
437472
engine.init_model(
438473
model_id=args.model_id,
439474
model_dir=str(model_dir),
440-
model_format="huggingface",
475+
model_format=_QWEN3_A8W8_FORMAT if args.model_format == _QWEN3_A8W8_FORMAT else "huggingface",
476+
decode_backend=args.decode_backend,
441477
runtime_config=RuntimeConfig(
442478
page_size=128,
443-
max_batch_size=16,
479+
max_batch_size=args.max_batch_size,
444480
max_seq_len=args.max_seq_len,
445481
max_new_tokens=args.max_new_tokens,
446482
device="cpu",
447-
kv_dtype="bfloat16",
448-
weight_dtype="float32",
483+
kv_dtype="int8" if args.model_format == _QWEN3_A8W8_FORMAT else "bfloat16",
484+
weight_dtype="bfloat16" if args.model_format == _QWEN3_A8W8_FORMAT else "float32",
449485
),
450486
)
451487
if collector is not None:

0 commit comments

Comments
 (0)