1010from __future__ import annotations
1111
1212import argparse
13+ import dataclasses
1314import statistics
1415import sys
1516import time
@@ -33,11 +34,17 @@ def _bootstrap_package_root() -> None:
3334
3435from python .core import GenerateConfig , LLMEngine , RuntimeConfig
3536from python .core .kv_cache import KvCacheManager
37+ from python .core .model_loader import ModelLoader
3638from python .core .parallel import ParallelConfig , parse_device_ids
3739from 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
3943from 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+
398425def 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