Faster generation
Three things make KwkDecoder produce more tokens per second without changing which tokens it produces: prompt
lookup (on by default), a small draft model, and batches. With a SmolLM2-135M draft, SmolLM2-1.7B generated about 2x
faster in int4 and 2.5x faster in bfloat16, with the same tokens.
Prompt lookup
generate(..., lookup=8) is the default: candidate tokens are copied from earlier in the text and checked in one pass,
so the output is exactly what the model would write alone. It pays most on text that repeats its input (summaries,
edits, code): about 2.4x on SmolLM2-360M. lookup=0 turns it off.
A draft model
A small model of the same family proposes tokens and the large one checks them in one pass. The draft must use the same tokenizer as the large model:
import torch
from transformers import LlamaConfig, LlamaForCausalLM
import kwker.cpu_backend as scb
from kwker.decode import KwkDecoder, runner_unsupported
def tiny(layers, hidden):
cfg = LlamaConfig(vocab_size=1000, hidden_size=hidden, intermediate_size=2 * hidden, num_hidden_layers=layers,
num_attention_heads=2, num_key_value_heads=1)
return LlamaForCausalLM(cfg).eval()
big, small = tiny(4, 256), tiny(1, 64) # in practice: from_pretrained(...) of a large and a small model
if scb.decoder_available() and runner_unsupported(big, 128) is None:
dec = KwkDecoder(big, max_cache_len=128, precision="int4")
draft = KwkDecoder(small, max_cache_len=128, precision="int4")
prompt = torch.tensor([[1, 5, 7, 9]])
fast = dec.generate(prompt, max_new_tokens=16, draft=draft)
print(torch.equal(fast, dec.generate(prompt, max_new_tokens=16)))
True
With sampling, a proposed token is kept with the probability the model gives it, so for a given seed the text is the
same with or without lookup and drafts. install(model, draft=small_model) does the same for Hugging Face's
generate.
Several sequences per step
One decode step reads every weight once, whether it serves one sequence or eight. KwkDecoder(..., max_batch=B) with
generate_batch(prompts, n) decodes several prompts together; for requests that arrive over time, use the serving
engine (Chat and serving).
Related
- Generate text with a Hugging Face model
- Compile with torch.compile: skipping guard checks in a compiled decode loop.