Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 29 additions & 3 deletions docs/DEPLOYMENT.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
25 changes: 24 additions & 1 deletion jevany/checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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):
Expand All @@ -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):
Expand All @@ -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):
Expand Down Expand Up @@ -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", "")
Expand Down Expand Up @@ -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
Expand Down
12 changes: 9 additions & 3 deletions jevany/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -79,13 +81,17 @@ 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()
options = replace(options, cuda_graphs=True,
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,
Expand All @@ -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))
Expand Down
Loading
Loading