diff --git a/docs/DEPLOYMENT.md b/docs/DEPLOYMENT.md index 42a6c00..493ff3b 100644 --- a/docs/DEPLOYMENT.md +++ b/docs/DEPLOYMENT.md @@ -265,6 +265,7 @@ exact checkpoint path unless their report says otherwise. | `JEVANY_COMPILE` | `0`, `default`, `reduce-overhead`, `max-autotune`, `max-autotune-no-cudagraphs` | `0` when direct graphs are enabled | `0` | Optional `torch.compile`; it cannot be combined with direct CUDA Graphs | | `JEVANY_CUDA_GRAPHS` / `--cuda-graphs` | `0`, `1` | `1` for single-question serving | `1` for single-question serving | Replays whole-backbone CUDA graphs on one GPU | | `JEVANY_CUDA_GRAPH_MAX_TOKENS` / `--cuda-graph-max-tokens` | positive integer | `2048` | `2048` | Longer rows stay on eager inference instead of being padded into a slower graph | +| `JEVANY_FUSED_KERNELS` / `--fused-kernels` | `0`, `1` | `1` with merged BF16 LoRA and direct graphs | `0` until measured | Fused Qwen3.5 kernels: one-kernel norms, merged projections and transposed weights | Qwen3.5 and Qwen3.8 mix full-attention layers with Gated DeltaNet layers. `JEVANY_ATTN` controls only the full-attention layers. Install @@ -374,10 +375,35 @@ changes, while 27B stayed at 207 with no argmax changes. Maximum probability changes were 0.0160 and 0.0146 respectively. These serving measurements do not replace the exact benchmark scores reported by the model cards. +`--fused-kernels` (`JEVANY_FUSED_KERNELS=1`, `LoadOptions(fused_kernels=True)`) +rewrites a merged Qwen3.5-architecture backbone (the Qwen3.5 and Qwen3.8 bases) +for batch-1 inference; see `jevany/fused_kernels.py`. Each RMSNorm runs as one +FLA kernel; the MLP gate and up projections, the four Gated DeltaNet input +projections and the attention query, key and value projections each run as one +GEMM; FLA's chunked kernel applies the DeltaNet gate and beta sigmoid itself; +and decoder weights are stored transposed so that cuBLAS reads both GEMM +operands in their natural layout. The original modules keep views of the fused +weights, so memory does not grow. The option needs the LoRA merged into the +backbone (`JEVANY_MERGE_BF16=1` for BF16 checkpoints; direct-token checkpoints +stay unmerged), one CUDA device and `flash-linear-attention`. It cannot be +combined with `JEVANY_COMPILE`, a device map or choice readout, and +`describe()` reports what was fused under `acceleration.fused_kernels`. + +On an A100-40GB with BF16 merged JevAny-Qwen3.5-4B, SDPA, FLA, causal-conv1d +and direct graphs, fused kernels changed public-JevBench forward latency (231 +requests × 3, `scripts/benchmark_latency.py --serving-kernels`) from 16.43 to +11.25 ms at the median, from 41.79 to 34.10 ms on average and from 130.14 to +112.30 ms at p90. Requests above the graph limit run eagerly and became 12% +faster as well. On the Transfer-v9 development suite with the same serving +settings, clean accuracy moved from 78.68% to 78.78% and NLL from 0.587 to +0.588; 5 of 1,264 argmax decisions changed and the largest probability change +was 0.044. The option has not been measured on 27B. + `scripts/benchmark_latency.py` accepts `--cuda-graphs`, -`--cuda-graph-max-tokens`, `--max-packed` and `--serving-kernels`. The last -option keeps the fused SDPA kernels used by `jevany serve` instead of selecting -the math kernel used for fp32-exact evaluation. +`--cuda-graph-max-tokens`, `--fused-kernels`, `--max-packed` and +`--serving-kernels`. The last option keeps the fused SDPA kernels used by +`jevany serve` instead of selecting the math kernel used for fp32-exact +evaluation. For a latency-oriented 27B deployment, keep compilation off and cap direct graphs at 2,048 tokens: diff --git a/jevany/checkpoint.py b/jevany/checkpoint.py index 43f6019..66be992 100644 --- a/jevany/checkpoint.py +++ b/jevany/checkpoint.py @@ -151,6 +151,11 @@ class LoadOptions: cuda_graph_max_tokens largest captured row. Longer rows use eager inference. The conservative default avoids padding regressions on long requests; raise it only after measuring the target model and workload. + fused_kernels + fused Qwen3.5 inference kernels (jevany.fused_kernels): one-kernel RMSNorms, merged projections, the + Gated DeltaNet gate inside FLA's kernel and transposed weight storage. Needs the LoRA merged, one + CUDA device and flash-linear-attention; the gain comes inside CUDA graphs. Probabilities are close + to, not identical with, the unfused path. """ dtype: torch.dtype | None = None merge: bool = True @@ -164,6 +169,7 @@ class LoadOptions: max_memory_gib: float | None = None cuda_graphs: bool = False cuda_graph_max_tokens: int = 2048 + fused_kernels: bool = False def __post_init__(self): if self.compile_mode not in (None, *COMPILE_MODES): @@ -181,6 +187,10 @@ def __post_init__(self): raise ValueError("cuda_graphs require the whole model on one GPU; disable device_map") if type(self.cuda_graph_max_tokens) is not int or self.cuda_graph_max_tokens < 1: raise ValueError("cuda_graph_max_tokens must be a positive integer") + if self.fused_kernels and self.compile_mode: + raise ValueError("fused_kernels and compile_mode cannot be combined") + if self.fused_kernels and self.device_map is not None: + raise ValueError("fused_kernels require the whole model on one GPU; disable device_map") @classmethod def from_env(cls, env=os.environ): @@ -195,7 +205,8 @@ def from_env(cls, env=os.environ): device_map=env.get("JEVANY_DEVICE_MAP") or None, max_memory_gib=float(env["JEVANY_MAX_MEMORY_GIB"]) if env.get("JEVANY_MAX_MEMORY_GIB") else None, cuda_graphs=env.get("JEVANY_CUDA_GRAPHS", "0") == "1", - cuda_graph_max_tokens=int(env.get("JEVANY_CUDA_GRAPH_MAX_TOKENS", "2048"))) + cuda_graph_max_tokens=int(env.get("JEVANY_CUDA_GRAPH_MAX_TOKENS", "2048")), + fused_kernels=env.get("JEVANY_FUSED_KERNELS", "0") == "1") def load_options_from_args(args, env=os.environ): @@ -248,6 +259,14 @@ def load(self, device, opts=LoadOptions()): # Keep that association intact instead of merging the wrapper away. merge = (merge and meta.decision_mode != "lm_token" and not adapter_config.get("trainable_token_indices")) + if opts.fused_kernels: + if not str(device).startswith("cuda"): + raise ValueError("fused_kernels run on CUDA only") + if not merge: + raise ValueError("fused_kernels need the LoRA merged into the backbone (JEVANY_MERGE_BF16=1 for BF16 " + "checkpoints); direct-token checkpoints and trainable token embeddings stay unmerged") + from .fused_kernels import load_kernels + load_kernels() # fail before loading weights when flash-linear-attention is missing saved_args = meta.extra.get("args", {}) lora_targets = saved_args.get("lora_targets", "all") explicit_targets = saved_args.get("lora_target_modules", "") @@ -285,7 +304,11 @@ def load(self, device, opts=LoadOptions()): "compile_mode": opts.compile_mode, "lora_merged": bool(merge), "approximate_bf16_merge": bool(merge and meta.weights_dtype == "bf16"), + "fused_kernels": None, } + if opts.fused_kernels: + from .fused_kernels import fuse_qwen3_5 + m.inference_acceleration["fused_kernels"] = fuse_qwen3_5(m.lm) if opts.compile_mode: m.lm.compile(mode=opts.compile_mode, fullgraph=False, dynamic=True) graphs = None diff --git a/jevany/cli.py b/jevany/cli.py index a04a0d7..ecea81b 100644 --- a/jevany/cli.py +++ b/jevany/cli.py @@ -60,6 +60,8 @@ def decide_main(argv: list[str]) -> None: parser.add_argument("--cuda-graphs", action="store_true", help="capture CUDA graphs for the local checkpoint") parser.add_argument("--cuda-graph-max-tokens", type=int, help="largest captured row; longer rows run eagerly (default 2048)") + parser.add_argument("--fused-kernels", action="store_true", + help="fused Qwen3.5 inference kernels for the local checkpoint (JEVANY_FUSED_KERNELS=1)") add_inference_arguments(parser) args = parser.parse_args(argv) try: @@ -79,6 +81,8 @@ def decide_main(argv: list[str]) -> None: ): parser.error("CUDA graph and native inference-limit flags cannot be used with --readout choice; " "use --choice-max-tokens") + if choice_options is not None and args.fused_kernels: + parser.error("--fused-kernels applies only to native readout") options = load_options_from_args(args) if args.cuda_graphs or args.cuda_graph_max_tokens is not None: options = options or LoadOptions.from_env() @@ -86,6 +90,8 @@ def decide_main(argv: list[str]) -> None: cuda_graph_max_tokens=(args.cuda_graph_max_tokens if args.cuda_graph_max_tokens is not None else options.cuda_graph_max_tokens)) + if args.fused_kernels: + options = replace(options or LoadOptions.from_env(), fused_kernels=True) settings = ({ "readout": "choice", "choice_temperature": choice_options.temperature, @@ -101,10 +107,10 @@ def decide_main(argv: list[str]) -> None: else: if (args.device or args.dtype or args.model_name is not None or args.readout != "native" or args.device_map is not None or args.max_memory_gib is not None or args.cuda_graphs - or args.cuda_graph_max_tokens is not None + or args.cuda_graph_max_tokens is not None or args.fused_kernels or any(getattr(args, item.name) is not None for item in fields(InferenceOptions))): - parser.error("device, dtype, placement, model-name, readout, cuda-graphs and inference limit options " - "require --checkpoint") + parser.error("device, dtype, placement, model-name, readout, cuda-graphs, fused-kernels and inference " + "limit options require --checkpoint") from .client import JevClient client = JevClient(args.base_url) print(json.dumps(client(request), indent=2, ensure_ascii=False)) diff --git a/jevany/fused_kernels.py b/jevany/fused_kernels.py new file mode 100644 index 0000000..65e413a --- /dev/null +++ b/jevany/fused_kernels.py @@ -0,0 +1,206 @@ +"""Fused Qwen3.5 inference kernels: ``LoadOptions(fused_kernels=True)``, ``JEVANY_FUSED_KERNELS=1`` or +``--fused-kernels``. + +At batch size one inside CUDA graphs (jevany.cudagraphs) a Qwen3.5 forward is a long chain of small kernels around +its GEMMs. ``fuse_qwen3_5`` rewrites a merged, eval-mode backbone in place: + +* every zero-centred RMSNorm runs as one FLA kernel, with ``1 + weight`` precomputed in fp32; +* each MLP computes its gate and up projections as one GEMM followed by FLA's fused SwiGLU; +* each Gated DeltaNet layer computes its four input projections as one GEMM. FLA's chunked delta-rule kernel takes + the raw gate and beta inputs and applies ``-exp(A_log) * softplus(a + dt_bias)``, the beta sigmoid, the q/k L2 + norm and the grouped value heads itself; the output norm is FLA's fused gated RMSNorm; +* each full-attention layer computes its query (with output gate), key and value projections as one GEMM; +* decoder weights are stored transposed, so cuBLAS reads both GEMM operands in their natural layout. + +The fused matrices replace the originals and the original ``nn.Linear`` modules keep views of them, so memory does +not grow. The Gated DeltaNet path runs only without a cache, attention mask or packed-sequence metadata and the +attention path only without a cache, which covers row-mode inference and CUDA-graph capture; other calls and training +take the original forwards. The +FLA kernels keep intermediates in fp32 where transformers rounds to BF16 between operations, so probabilities are +close to, not identical with, the unfused path. +""" +from __future__ import annotations + +import inspect +from importlib.metadata import PackageNotFoundError, version +from types import SimpleNamespace + +import torch +from torch import nn + +MIN_FLA = "0.5.0" # grouped value heads, and the gate computed inside the chunked delta-rule kernel + + +def _fla_version(): + for name in ("fla-core", "flash-linear-attention"): + try: + return version(name) + except PackageNotFoundError: + continue + return None + + +def load_kernels(): + """The flash-linear-attention functions the fused forwards call; ValueError when they are unavailable.""" + from packaging.version import InvalidVersion, Version + found = _fla_version() + try: + recent = found is not None and Version(found) >= Version(MIN_FLA) + except InvalidVersion: + recent = False + if not recent: + raise ValueError(f"fused_kernels need flash-linear-attention>={MIN_FLA} (pip install 'jevany[fast]'); " + f"found {found}") + from fla.modules.activations import swiglu + from fla.modules.fused_norm_gate import rms_norm_gated + from fla.modules.layernorm import rms_norm + from fla.ops.gated_delta_rule import chunk_gated_delta_rule + return SimpleNamespace( + version=found, swiglu=swiglu, rms_norm=rms_norm, rms_norm_gated=rms_norm_gated, + chunk_gated_delta_rule=chunk_gated_delta_rule, + # flash-linear-attention 0.5.0 has no in-kernel beta sigmoid. + beta_sigmoid_in_kernel="use_beta_sigmoid_in_kernel" in inspect.signature(chunk_gated_delta_rule).parameters, + ) + + +def _fuse(linears): + """One contiguous (in x total out) matrix for bias-free linears; each linear keeps a view of its columns.""" + fused = torch.cat([linear.weight.detach().t() for linear in linears], dim=1).contiguous() + start = 0 + for linear in linears: + linear.weight = nn.Parameter(fused[:, start:start + linear.out_features].t(), requires_grad=False) + start += linear.out_features + return fused + + +def _transpose_storage(linear): + """Store the (out x in) weight as the transpose of a contiguous (in x out) matrix: F.linear then computes + x @ W^T with no transposed GEMM operand, which cuBLAS runs faster for these small-row shapes.""" + linear.weight = nn.Parameter(linear.weight.detach().t().contiguous().t(), requires_grad=False) + + +def _replace_forward(module, fused_forward, kernels): + original = module.forward + + def forward(*args, **kwargs): + return fused_forward(module, original, kernels, *args, **kwargs) + + module.forward = forward + + +def _rms_norm(self, original, kernels, hidden_states): + if self.training: + return original(hidden_states) + return kernels.rms_norm(hidden_states, self.fused_weight, None, eps=self.eps) + + +def _mlp(self, original, kernels, x): + if self.training: + return original(x) + gate_up = x @ self.fused_weight + size = self.intermediate_size + return self.down_proj(kernels.swiglu(gate_up[..., :size], gate_up[..., size:])) + + +def _gated_deltanet(self, original, kernels, hidden_states, cache_params=None, attention_mask=None, **kwargs): + if (self.training or cache_params is not None or attention_mask is not None + or kwargs.get("cu_seq_lens_q") is not None): + return original(hidden_states, cache_params=cache_params, attention_mask=attention_mask, **kwargs) + batch, length, _ = hidden_states.shape + projected = hidden_states @ self.fused_weight # [qkv | z | b | a] + qkv_end = 2 * self.key_dim + self.value_dim + z_end = qkv_end + self.value_dim + b_end = z_end + self.num_v_heads + mixed = projected[..., :qkv_end].transpose(1, 2) + if mixed.stride(0) % 8 or mixed.stride(2) % 8: + mixed = mixed.contiguous() # causal-conv1d reads channel-last input only with 16-byte aligned rows + mixed = kernels.causal_conv1d(mixed, self.conv1d.weight.squeeze(1), self.conv1d.bias, activation=self.activation) + query, key, value = torch.split(mixed.transpose(1, 2), [self.key_dim, self.key_dim, self.value_dim], dim=-1) + beta = projected[..., z_end:b_end] + in_kernel = {"use_beta_sigmoid_in_kernel": True} if kernels.beta_sigmoid_in_kernel else {} + core, _ = kernels.chunk_gated_delta_rule( + query.reshape(batch, length, -1, self.head_k_dim), key.reshape(batch, length, -1, self.head_k_dim), + value.reshape(batch, length, -1, self.head_v_dim), g=projected[..., b_end:], + beta=beta if in_kernel else beta.sigmoid(), use_qk_l2norm_in_kernel=True, use_gate_in_kernel=True, + A_log=self.A_log, dt_bias=self.dt_bias, output_final_state=False, **in_kernel) + core = kernels.rms_norm_gated(core.reshape(-1, self.head_v_dim), + projected[..., qkv_end:z_end].reshape(-1, self.head_v_dim), + self.norm.weight, None, activation="swish", eps=self.norm.variance_epsilon) + return self.out_proj(core.reshape(batch, length, -1)) + + +def _attention(self, original, kernels, hidden_states, position_embeddings, attention_mask, past_key_values=None, + **kwargs): + if self.training or past_key_values is not None: + return original(hidden_states, position_embeddings, attention_mask, past_key_values=past_key_values, + **kwargs) + qwen = kernels.qwen + input_shape = hidden_states.shape[:-1] + hidden_shape = (*input_shape, -1, self.head_dim) + projected = hidden_states @ self.fused_weight # [query and gate | key | value] + query_end, key_end = self.fused_sizes + query, gate = torch.chunk(projected[..., :query_end].view(*input_shape, -1, self.head_dim * 2), 2, dim=-1) + query = self.q_norm(query.reshape(hidden_shape)).transpose(1, 2) + key = self.k_norm(projected[..., query_end:key_end].reshape(hidden_shape)).transpose(1, 2) + value = projected[..., key_end:].reshape(hidden_shape).transpose(1, 2) + cos, sin = position_embeddings + query, key = qwen.apply_rotary_pos_emb(query, key, cos, sin) + attention = qwen.ALL_ATTENTION_FUNCTIONS.get_interface(self.config._attn_implementation, + qwen.eager_attention_forward) + output, weights = attention(self, query, key, value, attention_mask, dropout=0.0, scaling=self.scaling, **kwargs) + output = output.reshape(*input_shape, -1).contiguous() * torch.sigmoid(gate.reshape(*input_shape, -1)) + return self.o_proj(output), weights + + +def fuse_qwen3_5(lm): + """Rewrite the Qwen3.5 modules of ``lm`` in place for fused inference and report what changed. + + ``lm`` must be a backbone with its LoRA merged, on one CUDA device, with its final weights loaded.""" + from transformers.models.qwen3_5 import modeling_qwen3_5 as qwen + modules = list(lm.modules()) + layers = [module for module in modules if isinstance(module, qwen.Qwen3_5DecoderLayer)] + if not layers: + raise ValueError("fused_kernels support Qwen3.5 backbones only") + devices = {parameter.device for parameter in lm.parameters()} + if len(devices) != 1 or next(iter(devices)).type != "cuda": + raise ValueError("fused_kernels need the whole backbone on one CUDA device") + if any(hasattr(module, "lora_A") for module in modules): + raise ValueError("fused_kernels need the LoRA merged into the backbone") + kernels = load_kernels() + kernels.qwen, kernels.causal_conv1d = qwen, qwen.causal_conv1d_fn + counts = dict.fromkeys(("rms_norm", "mlp", "gated_deltanet", "attention", "transposed_linear"), 0) + + def bias_free(*linears): + return all(linear.bias is None for linear in linears) + + with torch.no_grad(): + for module in modules: + if isinstance(module, qwen.Qwen3_5RMSNorm): + module.fused_weight = (1.0 + module.weight.detach().float()).contiguous() + _replace_forward(module, _rms_norm, kernels) + counts["rms_norm"] += 1 + elif (isinstance(module, qwen.Qwen3_5MLP) and module.config.hidden_act in ("silu", "swish") + and bias_free(module.gate_proj, module.up_proj)): + module.fused_weight = _fuse([module.gate_proj, module.up_proj]) + _replace_forward(module, _mlp, kernels) + counts["mlp"] += 1 + elif isinstance(module, qwen.Qwen3_5GatedDeltaNet) and module.activation in ("silu", "swish"): + projections = [module.in_proj_qkv, module.in_proj_z, module.in_proj_b, module.in_proj_a] + if bias_free(*projections): + module.fused_weight = _fuse(projections) + _replace_forward(module, _gated_deltanet, kernels) + counts["gated_deltanet"] += 1 + elif (isinstance(module, qwen.Qwen3_5Attention) + and bias_free(module.q_proj, module.k_proj, module.v_proj)): + query, key = module.q_proj.out_features, module.k_proj.out_features + module.fused_weight = _fuse([module.q_proj, module.k_proj, module.v_proj]) + module.fused_sizes = (query, query + key) + _replace_forward(module, _attention, kernels) + counts["attention"] += 1 + for layer in layers: + for linear in layer.modules(): + if isinstance(linear, nn.Linear) and linear.weight.is_contiguous(): + _transpose_storage(linear) + counts["transposed_linear"] += 1 + torch.cuda.empty_cache() + return {"backbone": "qwen3_5", "flash_linear_attention": kernels.version, **counts} diff --git a/jevany/letter_predictor.py b/jevany/letter_predictor.py index ea9c656..151d194 100644 --- a/jevany/letter_predictor.py +++ b/jevany/letter_predictor.py @@ -300,6 +300,8 @@ def __init__( raise ValueError("native-head temperature is not used by choice readout; use temperature") if self.options.cuda_graphs: raise ValueError("CUDA graph capture is available only for native readout") + if self.options.fused_kernels: + raise ValueError("fused kernels are available only for native readout") if self.exact_cuda_kernels_applied: # Match LocalPredictor: benchmark evaluation disables approximate # TF32 and fused SDPA kernels, while exact_kernels=False retains diff --git a/jevany/runtime.py b/jevany/runtime.py index 1d1205e..fbbca6d 100644 --- a/jevany/runtime.py +++ b/jevany/runtime.py @@ -213,6 +213,8 @@ def from_pretrained( options = replace(options, attn="sdpa") if readout == "choice" and options.cuda_graphs: raise ValueError("CUDA graph capture is available only for native readout") + if readout == "choice" and options.fused_kernels: + raise ValueError("fused kernels are available only for native readout") if readout == "choice" and options.temperature is not None: raise ValueError("JEVANY_TEMPERATURE applies to the native head; use choice_temperature") predictor = None diff --git a/jevany/serve.py b/jevany/serve.py index 49944ff..0710ffe 100644 --- a/jevany/serve.py +++ b/jevany/serve.py @@ -226,6 +226,9 @@ def main(argv=None): help="capture CUDA graphs at startup (row-mode backbones on one GPU); same as JEVANY_CUDA_GRAPHS=1") ap.add_argument("--cuda-graph-max-tokens", type=int, help="largest captured row; longer rows run eagerly (default 2048)") + ap.add_argument("--fused-kernels", action="store_true", + help="fused Qwen3.5 inference kernels (merged LoRA, one GPU, flash-linear-attention); " + "same as JEVANY_FUSED_KERNELS=1") add_inference_arguments(ap) ap.add_argument("--host", default="127.0.0.1") ap.add_argument("--port", type=int, default=8008) @@ -240,12 +243,16 @@ def main(argv=None): ): ap.error("CUDA graph and native inference-limit flags cannot be used with --readout choice; " "use --choice-max-tokens") + if choice_options is not None and a.fused_kernels: + ap.error("--fused-kernels applies only to native readout") options = load_options_from_args(a) if a.cuda_graphs or a.cuda_graph_max_tokens is not None: options = options or LoadOptions.from_env() options = replace(options, cuda_graphs=True, cuda_graph_max_tokens=(a.cuda_graph_max_tokens if a.cuda_graph_max_tokens is not None else options.cuda_graph_max_tokens)) + if a.fused_kernels: + options = replace(options or LoadOptions.from_env(), fused_kernels=True) application = create_app(a.run, device=a.device, dtype=a.dtype, model_name=a.model_name, options=options, inference_options=(inference_options_from_args(a) diff --git a/scripts/benchmark_latency.py b/scripts/benchmark_latency.py index 526c6a2..9df40b6 100644 --- a/scripts/benchmark_latency.py +++ b/scripts/benchmark_latency.py @@ -96,6 +96,8 @@ def main(argv=None): help="replay captured CUDA graphs (LoadOptions.cuda_graphs; also JEVANY_CUDA_GRAPHS=1)") parser.add_argument("--cuda-graph-max-tokens", type=int, help="largest captured row; longer rows run eagerly (default 2048)") + parser.add_argument("--fused-kernels", action="store_true", + help="fused Qwen3.5 inference kernels (LoadOptions.fused_kernels; also JEVANY_FUSED_KERNELS=1)") parser.add_argument("--serving-kernels", action="store_true", help="keep PyTorch's fused SDPA kernels as `jevany serve` does, instead of the math kernel " "LocalPredictor selects for fp32-exact evaluation") @@ -108,6 +110,8 @@ def main(argv=None): parser.error("records must be nonnegative; warmup and repeats must be positive") if letter_options is not None and (args.cuda_graphs or args.cuda_graph_max_tokens is not None): parser.error("CUDA graphs apply only to native readout") + if letter_options is not None and args.fused_kernels: + parser.error("fused kernels apply only to native readout") target = Path(args.out) if target.exists(): parser.error("refusing to overwrite an existing latency report") @@ -132,6 +136,8 @@ def main(argv=None): cuda_graph_max_tokens=(args.cuda_graph_max_tokens if args.cuda_graph_max_tokens is not None else options.cuda_graph_max_tokens)) + if args.fused_kernels: + options = replace(options, fused_kernels=True) if letter_options is None: predictor = LocalPredictor(args.run, args.device, options, max_packed=args.max_packed, exact_kernels=not args.serving_kernels) @@ -176,6 +182,7 @@ def main(argv=None): "merge_bf16": options.merge_bf16, "compile_mode": options.compile_mode, "device_map": options.device_map, "max_memory_gib": options.max_memory_gib, "cuda_graph_max_tokens": options.cuda_graph_max_tokens, + "fused_kernels": options.fused_kernels, }, readout=args.readout, letter_readout=(dict(predictor.provenance) if letter_options is not None else None), @@ -183,6 +190,10 @@ def main(argv=None): getattr(predictor, "model", getattr(predictor, "pointer_model", None)), "inference_acceleration", {}, ).get("cuda_graphs"), + fused_kernels=getattr( + getattr(predictor, "model", getattr(predictor, "pointer_model", None)), + "inference_acceleration", {}, + ).get("fused_kernels"), ) target.parent.mkdir(parents=True, exist_ok=True) with target.open("x") as output: diff --git a/tests/test_fused_kernels.py b/tests/test_fused_kernels.py new file mode 100644 index 0000000..6838b85 --- /dev/null +++ b/tests/test_fused_kernels.py @@ -0,0 +1,132 @@ +"""Fused Qwen3.5 inference kernels: option plumbing and weight sharing on CPU; parity with the unfused path on CUDA +with flash-linear-attention installed.""" +import importlib.util + +import pytest +import torch +from torch import nn + +from jevany.checkpoint import Checkpoint, LoadOptions +from jevany.fused_kernels import _fuse, _transpose_storage, fuse_qwen3_5 +from jevany.model import DecisionModel, load_tokenizer +from test_backbones import make_base + +fla = pytest.mark.skipif(not torch.cuda.is_available() or importlib.util.find_spec("fla") is None, + reason="needs CUDA and flash-linear-attention") + + +def test_fused_kernel_option_from_env(): + assert not LoadOptions.from_env({}).fused_kernels + assert LoadOptions.from_env({"JEVANY_FUSED_KERNELS": "1"}).fused_kernels + assert not LoadOptions.from_env({"JEVANY_FUSED_KERNELS": "0"}).fused_kernels + with pytest.raises(ValueError, match="cannot be combined"): + LoadOptions(fused_kernels=True, compile_mode="default") + with pytest.raises(ValueError, match="one GPU"): + LoadOptions(fused_kernels=True, device_map="auto") + + +def test_fused_projections_share_one_transposed_matrix(): + torch.manual_seed(0) + linears = [nn.Linear(6, size, bias=False) for size in (4, 3, 5)] + x = torch.randn(2, 6) + expected = [linear(x) for linear in linears] + fused = _fuse(linears) + assert fused.shape == (6, 12) and fused.is_contiguous() + torch.testing.assert_close(x @ fused, torch.cat(expected, dim=-1)) + for linear, reference in zip(linears, expected): + assert linear.weight.shape == (linear.out_features, 6) + assert linear.weight.untyped_storage().data_ptr() == fused.untyped_storage().data_ptr() + torch.testing.assert_close(linear(x), reference) + + +def test_transposed_storage_keeps_linear_outputs(): + torch.manual_seed(0) + linear, x = nn.Linear(6, 4, bias=False), torch.randn(3, 6) + expected = linear(x) + _transpose_storage(linear) + assert linear.weight.shape == (4, 6) and linear.weight.t().is_contiguous() + torch.testing.assert_close(linear(x), expected) + + +@pytest.mark.parametrize("family,message", [("qwen35", "one CUDA device"), ("llama", "Qwen3.5")]) +def test_fusion_checks_backbone_and_device(tmp_path, family, message): + base = tmp_path / "base" + make_base(base, family, legacy=True) + model = DecisionModel(base, load_tokenizer(base), "cpu", lora=2, head_dim=8) + with pytest.raises(ValueError, match=message): + fuse_qwen3_5(model.lm) + + +def test_checkpoint_and_choice_readout_reject_unsupported_fused_kernels(tmp_path): + from jevany import JevModel + from test_serving import make_checkpoint + checkpoint = make_checkpoint(tmp_path, "qwen35", legacy=True) + with pytest.raises(ValueError, match="CUDA only"): + Checkpoint(checkpoint).load("cpu", LoadOptions(fused_kernels=True)) + with pytest.raises(ValueError, match="native readout"): + JevModel.from_pretrained(checkpoint, device="cpu", options=LoadOptions(fused_kernels=True), readout="choice") + + +def test_remote_decide_rejects_fused_kernels(tmp_path, capsys): + from jevany import Choice, SystemOneRequest + from jevany.cli import main + request = tmp_path / "request.json" + request.write_text(SystemOneRequest(state="state", questions={ + "q": Choice(instructions="choose", criteria={"a": None, "b": None})}).model_dump_json()) + with pytest.raises(SystemExit): + main(["decide", str(request), "--fused-kernels"]) + assert "require --checkpoint" in capsys.readouterr().err + + +def _qwen35_backbone(dtype): + from transformers import Qwen3_5TextConfig + from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5TextModel + config = Qwen3_5TextConfig( + vocab_size=128, hidden_size=256, intermediate_size=512, num_hidden_layers=4, num_attention_heads=4, + num_key_value_heads=2, head_dim=64, layer_types=["linear_attention"] * 3 + ["full_attention"], + linear_num_key_heads=2, linear_num_value_heads=4, linear_key_head_dim=64, linear_value_head_dim=64, + max_position_embeddings=1024) + config._attn_implementation = "sdpa" + torch.manual_seed(0) + model = Qwen3_5TextModel(config) + with torch.no_grad(): # norm weights start at their identity values; exercise the scales too + for module in model.modules(): + if type(module).__name__ in ("Qwen3_5RMSNorm", "Qwen3_5RMSNormGated"): + module.weight.add_(0.1 * torch.randn_like(module.weight)) + return model.to("cuda", dtype).eval() + + +@fla +@pytest.mark.parametrize("dtype,tolerance", [(torch.float32, 1e-3), (torch.bfloat16, 3e-2)]) +def test_fused_backbone_matches_transformers(dtype, tolerance): + model = _qwen35_backbone(dtype) + ids = torch.randint(0, 128, (1, 200), device="cuda") + positions = torch.arange(200, device="cuda")[None] + with torch.no_grad(): + expected = model(input_ids=ids, position_ids=positions, use_cache=False).last_hidden_state.float() + counts = fuse_qwen3_5(model) + actual = model(input_ids=ids, position_ids=positions, use_cache=False).last_hidden_state.float() + assert {key: counts[key] for key in ("rms_norm", "mlp", "gated_deltanet", "attention")} == { + "rms_norm": 11, "mlp": 4, "gated_deltanet": 3, "attention": 1} + error = ((actual - expected).norm() / expected.norm()).item() + print(f"relative error {dtype}: {error:.2e}") + assert error < tolerance + + +@fla +def test_checkpoint_fused_kernels_replay_and_report(tmp_path): + from jevany import Choice, JevModel, SystemOneRequest + from test_serving import make_checkpoint + checkpoint = make_checkpoint(tmp_path, "qwen35", legacy=True) + request = SystemOneRequest(state="state " * 20, questions={ + "choice": Choice(instructions="choose", criteria={"a": None, "b": None})}) + plain = JevModel.from_pretrained(checkpoint, device="cuda", options=LoadOptions(cuda_graphs=True)) + fused = JevModel.from_pretrained(checkpoint, device="cuda", + options=LoadOptions(cuda_graphs=True, fused_kernels=True)) + expected, actual = plain(request), fused(request) + assert actual["answers"]["choice"]["probabilities"] == pytest.approx( + expected["answers"]["choice"]["probabilities"], abs=1e-3) + acceleration = fused.describe()["acceleration"] + assert acceleration["fused_kernels"]["gated_deltanet"] == 1 and acceleration["fused_kernels"]["attention"] == 1 + assert acceleration["cuda_graphs"]["graph_calls"] == 1 + assert plain.describe()["acceleration"]["fused_kernels"] is None