Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

43. Beyond text: images, embeddings, reranking and many LoRAs

In this chapter

  • Turn an image into decoder input rows with a small vision transformer and projector.
  • Give image patches three-dimensional rotary positions, while keeping the scheduler's physical token positions.
  • Serve embeddings and query-document scores without an autoregressive loop.
  • Run requests with different LoRA adapters in the same batch, using adapter slots and shrink/expand kernels.
  • Make prefix reuse depend on everything that changes activations, including images and adapter weights.

You will build

VisionEncoder.forward, expand_images and mrope_cos_sin in engine/multimodal.py; pool_hidden in engine/pooling.py; grouped_lora and AdapterBank.load in engine/serve/adapters.py; ModelRunner.prepare_features; and the BGMV/SGMV kernels.

Time: 6-8 hours. GPU: optional; the kernels also run in Triton's interpreter.

What changes at the input and output

The decoder sees a matrix of hidden vectors, however those vectors were obtained. A text embedding lookup is one source. A vision encoder is another. The final hidden vectors can go to an LM head to generate tokens, to a pooling operation to produce an embedding, or to a trained scalar head to score a query-document pair.

LoRA changes something different: which function a linear layer computes for each row. Chapter 22 attached one adapter to a model. A service needs to keep a base model resident while requests choose different adapters, without splitting the batch into one base-model forward per adapter.

requestinput boundarywork in the modeloutput boundary
text generationtoken IDscausal decoder, repeated decodeincremental text
image + texttext IDs and projected image patchesthe same causal decoderincremental text
embeddingstoken IDs and a padding maskone encoder/decoder forwardpooled vector
rerankingtokenized query-document pairsone encoder forward and score headscalar scores
any supported LoRA generationtoken IDs and adapter identitybase projections plus per-row low-rank updatesincremental text

Images and adapters belong in Chapter 31’s request and batch metadata. Pooling belongs beside generation: it needs one bounded forward, not KV allocation and hundreds of decode steps.

An image becomes tokens

Take a 16 × 16 RGB image and divide it into 4 × 4 patches. There are 16 patches; each contains 48 pixel values. A convolution with kernel and stride 4 is a learned linear projection of those values into the vision width. Add a learned position vector to every patch, then run transformer layers with bidirectional attention: every patch may see every other patch. Finally, an MLP projects each patch into the decoder’s hidden width.

RGB pixels [B,3,16,16]
       │ patch projection + learned positions
       ▼
16 patch rows [B,16,vision_width]
       │ bidirectional transformer + normalization
       │ trained projector
       ▼
decoder rows [B,16,hidden_size]
class VisionEncoder(nn.Module):
    def __init__(self, config=None):
        super().__init__()
        self.cfg = c = config or VisionConfig()
        if c.image_size % c.patch_size or c.width % c.heads:
            raise ValueError("Image size must divide into patches, width into heads")
        self.patch = nn.Conv2d(3, c.width, c.patch_size, stride=c.patch_size)
        self.position = nn.Parameter(torch.randn((c.image_size // c.patch_size) ** 2, c.width) * 0.02)
        self.blocks = nn.ModuleList(nn.TransformerEncoderLayer(
            c.width, c.heads, 4 * c.width, dropout=0, activation="gelu", batch_first=True, norm_first=True)
            for _ in range(c.layers))
        self.norm = nn.LayerNorm(c.width)
        self.projector = nn.Sequential(nn.Linear(c.width, c.output_width), nn.GELU(),
                                       nn.Linear(c.output_width, c.output_width))

    def forward(self, pixels):
        """Normalized RGB [B,3,H,W] -> projected patch rows [B,P,D].  (Your engine: Chapter 43)"""
        if pixels.shape[1:] != (3, self.cfg.image_size, self.cfg.image_size):
            raise ValueError("Pixels must match the configured RGB image size")
        x = self.patch(pixels).flatten(2).transpose(1, 2) + self.position
        for block in self.blocks:
            x = block(x)                       # bidirectional attention within the image
        return self.projector(self.norm(x))

There is no magic in the projector. It has to be trained so that its outputs mean something to the language model. Random weights are enough to test shapes and serving correctness; they do not make an image assistant. Real VLMs may resize into several crops, merge groups of patches, use a different vision tower or insert boundary tokens. Loading a Qwen-VL checkpoint requires all those exact conventions, alongside its weights.

At the text boundary, a prompt contains one reserved <image> placeholder for each image. The processor expands each placeholder into as many reserved IDs as there are projected patches, and records which absolute positions get which vectors:

@dataclass
class VisionPrompt:
    token_ids: list
    image_rows: dict                 # absolute token position -> projected embedding (CPU)
    positions: torch.Tensor         # [3,T]: time, height, width
    delta: int                      # next text position minus physical prompt length
    sections: tuple

    def features(self):
        return {"image_rows": self.image_rows, "mrope_positions": self.positions,
                "mrope_delta": self.delta, "mrope_sections": self.sections}


def expand_images(ids, placeholder_id, embeddings, grids, sections):
    """One placeholder becomes T*H*W patch rows; text shares all three axes.
    Resume text at max(image coordinates)+1, not after the number of patches.  (Your engine: Chapter 43)
    """
    if ids.count(placeholder_id) != len(embeddings) or len(grids) != len(embeddings):
        raise ValueError("Each image needs exactly one placeholder and one grid")
    tokens, axes, rows, cursor, index = [], [], {}, 0, 0
    for token in ids:
        if token != placeholder_id:
            tokens.append(token), axes.append((cursor, cursor, cursor))
            cursor += 1
            continue
        grid, features = grids[index], embeddings[index].detach().cpu().clone()
        t, h, w = grid
        if min(grid) < 1 or features.ndim != 2 or len(features) != t * h * w:
            raise ValueError("Projected image length does not match its grid")
        for tt in range(t):
            for hh in range(h):
                for ww in range(w):
                    rows[len(tokens)] = features[(tt * h + hh) * w + ww]
                    tokens.append(placeholder_id), axes.append((cursor + tt, cursor + hh, cursor + ww))
        cursor += max(grid)
        index += 1
    return VisionPrompt(tokens, rows, torch.tensor(axes, dtype=torch.long).T.contiguous(),
                        cursor - len(tokens), tuple(sections))

The reserved IDs give the scheduler a physical sequence to count and page. Their embedding lookup is replaced by the projected rows before the decoder runs. A 16-patch image consumes 16 KV positions, even though it began as one placeholder. Expand before validating prompt length and reserving KV memory.

VisionProcessor supplies a fixed-resolution demo processor. It accepts base64 image data URIs, checks compressed bytes and pixel count, converts to RGB and normalizes to [-1, 1]. Its decoder inputs are detached CPU tensors, which can cross Chapter 36’s process queues. Install Pillow for this boundary (uv pip install pillow); tensor-only milestones don’t need it. The service does not fetch remote URLs. A deployed processor needs a controlled fetcher and the model’s trained preprocessing recipe if URLs are part of its API.

Two meanings of position

Ordinary RoPE has one coordinate, the text position. Image patches have spatial coordinates; videos also have time. M-RoPE assigns different frequency pairs to time, height and width. With a 12-wide head there are 6 rotary pairs. Sections (2, 2, 2) assign 2 pairs to each axis; the selected angles are then repeated for the head’s split-half pairing (Chapter 17).

For text, all three coordinates are equal. M-RoPE then reduces exactly to ordinary RoPE. For an image, flattening patch rows should not erase the fact that neighbouring rows may be vertically or horizontally adjacent:

physical row in A <image> Bcontenttimeheightwidth
0A000
1image patch (0,0)111
2image patch (0,1)112
3image patch (1,0)121
4image patch (1,1)122
5B333

The next generated token has physical position 6 but rotary position 4. Store the continuation delta, here -2, and compute text rotary positions as physical position + delta after the prompt. Do not change slot_mapping, causal masking or sequence lengths: those still use the physical order.

def mrope_cos_sin(positions, head_dim, theta, sections, dtype):
    """positions [3,N], sections count frequency pairs on each axis.  (Your engine: Chapter 43)
    Split-half RoPE repeats the chosen angles for both halves of the head.
    """
    if positions.ndim != 2 or positions.shape[0] != 3 or len(sections) != 3 \
            or min(sections) < 0 or sum(sections) * 2 != head_dim:
        raise ValueError("M-RoPE needs three sections totaling head_dim/2 pairs")
    inv = theta ** (-torch.arange(0, head_dim, 2, device=positions.device).float() / head_dim)
    selected, start = [], 0
    for axis, count in enumerate(sections):
        selected.append(positions[axis, :, None].float() * inv[start:start + count])
        start += count
    angles = torch.cat(selected, -1)
    angles = torch.cat((angles, angles), -1)[:, None]
    return angles.cos().to(dtype), angles.sin().to(dtype)

This implements the position mechanism described by Qwen2-VL, with a deliberately small image processor. The model must choose the sections; the engine fixes them on first use if they are not configured, then rejects mismatches. This path uses ordinary rotary frequencies and refuses scaled RoPE rather than silently mixing the two recipes.

Chunked prefill, cache hits and recomputation

An image can span several prefill chunks. Storing projected rows by absolute prompt position makes that case simple: when a request computes positions c:c+n, replace only image rows inside that interval. A cached prefix skips rows already computed; recompute preemption reads the same saved vectors again. M-RoPE coordinates are sliced the same way, with the continuation delta used during decode.

    @torch.inference_mode()
    def prepare_features(self, batch, scheduled):
        """Map absolute feature positions and adapter slots to this step's ragged rows.  (Your engine: Chapter 43)
        A prefill chunk can cut through an image; recomputation uses the same saved rows.
        """
        extras = [r.extra for r, _ in scheduled]
        if any("lora_slot" in e for e in extras):
            slots = torch.full_like(batch.input_ids, -1)
            for (request, n), start in zip(scheduled, batch.meta.query_start_loc_cpu):
                slots[start:start + n] = request.extra.get("lora_slot", -1)
            batch.extra["lora_ids"] = slots
        mrope = [e for e in extras if "mrope_positions" in e]
        if mrope:
            sections = mrope[0]["mrope_sections"]
            if any(e["mrope_sections"] != sections for e in mrope):
                raise ValueError("All requests must use the model's M-RoPE sections")
            axes = batch.positions[None].expand(3, -1).clone()
            for (request, n), start in zip(scheduled, batch.meta.query_start_loc_cpu):
                if "mrope_positions" not in request.extra:
                    continue
                for offset in range(n):
                    pos = request.num_computed_tokens + offset
                    if pos < request.num_prompt_tokens:
                        axes[:, start + offset] = request.extra["mrope_positions"][:, pos].to(self.device)
                    else:
                        axes[:, start + offset] = pos + request.extra["mrope_delta"]
            batch.extra.update(mrope_positions=axes, mrope_sections=sections)
        replacements = []
        for (request, n), start in zip(scheduled, batch.meta.query_start_loc_cpu):
            for pos, row in request.extra.get("image_rows", {}).items():
                offset = pos - request.num_computed_tokens
                if 0 <= offset < n:
                    replacements.append((start + offset, row))
        if replacements:
            embeds = self.flat.backbone.embed_tokens(batch.input_ids)
            for row, value in replacements:
                embeds[row] = value.to(embeds)
            batch.extra["embeds"] = embeds

Two requests can have identical placeholder IDs and entirely different images. Chapter 31’s token hash alone would alias their KV. validate_features snapshots projected rows and derives hashes from their actual values, positions, sections and continuation delta. Those hashes are added to Request.cache_key, so the block hash chain identifies the inputs the decoder actually consumed. The caller cannot assert that two images share a hash. This conservative implementation salts every block with all image inputs, including text blocks before an image; a more selective cache can begin salting at the first affected block.

Images are encoded in the frontend, then their CPU rows travel through AsyncLLM.generate(..., features=...). For a vision-capable Server, set multimodal_processor=VisionProcessor(...) and ensure the tokenizer recognizes its reserved marker. A text-only server rejects image parts. Dropping them and answering only the text would look successful while computing the wrong request.

Embeddings: one forward, one vector

An embedding model is trained to put related inputs near each other. Pooling determines which vector represents the sequence:

  • Mean: sum the visible token vectors and divide by their count.
  • Last: select the last visible token, useful for decoder embedding models.
  • CLS: select the designated first token, useful for encoders trained with that convention.

For cosine similarity, normalize the pooled vector to unit length. Then the dot product equals cosine similarity. The pooling mode, query/document prefixes and normalization must match training; a random causal decoder with mean pooling proves the plumbing, not retrieval quality.

def pool_hidden(hidden, mask, mode="mean", normalize=True):
    """Exclude padding, select last/CLS or average, optionally unit-normalize.  (Your engine: Chapter 43)"""
    if hidden.ndim != 3 or mask.shape != hidden.shape[:2] or not mask.any(-1).all():
        raise ValueError("Each sequence needs at least one visible token")
    if mode == "mean":
        result = (hidden.float() * mask[..., None]).sum(1) / mask.sum(1, keepdim=True)
    elif mode == "last":
        index = torch.arange(mask.shape[1], device=mask.device).expand_as(mask).masked_fill(~mask.bool(), -1).max(1).values
        result = hidden[torch.arange(len(hidden), device=hidden.device), index].float()
    elif mode == "cls":
        if not mask[:, 0].all():
            raise ValueError("CLS pooling requires an unpadded first token")
        result = hidden[:, 0].float()
    else:
        raise ValueError("Pooling mode must be mean|last|cls")
    return F.normalize(result, dim=-1) if normalize else result

Padding must contribute neither vectors nor a denominator. The mean of [h0, h1, padding] is (h0+h1)/2, not /3. Last pooling finds the last visible index, which also works with left padding. The supplied DecoderEncoder uses right padding and causal attention, so valid tokens cannot see subsequent padding. HFEncoder instead passes the attention mask to a bidirectional encoder.

class PoolingService:
    def __init__(self, encoder, tokenizer, model_name, mode="mean", normalize=True, score_head=None,
                 max_model_len=512, max_batch=32, pad_id=0):
        self.encoder, self.tokenizer, self.model_name = encoder.eval(), tokenizer, model_name
        self.mode, self.normalize, self.score_head = mode, normalize, score_head
        self.max_model_len, self.max_batch, self.pad_id = max_model_len, max_batch, pad_id
        self.lock = asyncio.Lock()              # bound device work; don't overlap calls on one encoder

    def inputs(self, value):
        if isinstance(value, str) or isinstance(value, list) and value and all(type(x) is int for x in value):
            value = [value]
        if not isinstance(value, list) or not 1 <= len(value) <= self.max_batch:
            raise ValueError("Input must be a nonempty bounded batch")
        ids = []
        for item in value:
            row = self.tokenizer.encode(item) if isinstance(item, str) else item
            if not isinstance(row, list) or not row or not all(type(x) is int and x >= 0 for x in row):
                raise ValueError("Each input must be text or nonempty token ids")
            if len(row) > self.max_model_len:
                raise ValueError("Embedding input exceeds the model context limit")
            ids.append(row)
        return ids

    @torch.inference_mode()
    def encode(self, ids):
        device = next(self.encoder.parameters()).device
        tokens = torch.full((len(ids), max(map(len, ids))), self.pad_id, dtype=torch.long, device=device)
        mask = torch.zeros_like(tokens, dtype=torch.bool)
        for i, row in enumerate(ids):
            tokens[i, :len(row)] = torch.tensor(row, device=device)
            mask[i, :len(row)] = True
        hidden = self.encoder(tokens, mask)
        pooled = pool_hidden(hidden, mask, self.mode, self.normalize)
        if self.score_head is not None:
            return self.score_head(pooled.to(next(self.score_head.parameters()))).float().flatten().cpu()
        return pooled.cpu()

    async def __call__(self, body):
        if self.score_head is not None:
            raise ValueError("This service is a reranker")
        if body.get("dimensions") is not None:
            raise ValueError("Dimension truncation needs a model trained for it")
        fmt = body.get("encoding_format", "float")
        if fmt not in ("float", "base64"):
            raise ValueError("encoding_format must be float|base64")
        ids = self.inputs(body.get("input"))
        async with self.lock:
            vectors = await asyncio.to_thread(self.encode, ids)
        embeddings = vectors.tolist() if fmt == "float" else [
            base64.b64encode(v.contiguous().numpy().astype("<f4").tobytes()).decode() for v in vectors]
        count = sum(map(len, ids))
        return {"object": "list", "model": self.model_name,
                "data": [{"object": "embedding", "index": i, "embedding": v} for i, v in enumerate(embeddings)],
                "usage": {"prompt_tokens": count, "total_tokens": count}}

    async def rerank(self, body):
        if self.score_head is None:
            raise ValueError("Reranking needs a trained scalar score head")
        query, docs = body.get("query"), body.get("documents")
        if not isinstance(query, str) or not isinstance(docs, list) or not docs or not all(isinstance(d, str) for d in docs):
            raise ValueError("Reranking needs a query and a nonempty list of text documents")
        # The tokenizer owns the trained pair convention; do not invent a separator.
        pair = getattr(self.tokenizer, "encode_pair", None)
        if pair is None:
            raise ValueError("Reranking requires a tokenizer with encode_pair(query, document)")
        ids = self.inputs([pair(query, doc) for doc in docs])
        top_n = body.get("top_n", len(docs))
        if type(top_n) is not int or not 1 <= top_n <= len(docs):
            raise ValueError("top_n must be between one and the number of documents")
        async with self.lock:
            scores = await asyncio.to_thread(self.encode, ids)
        if scores.shape != (len(docs),):
            raise ValueError("Score head must return one scalar per query-document pair")
        order = sorted(range(len(docs)), key=lambda i: (-float(scores[i]), i))[:top_n]
        return {"model": self.model_name, "results": [{"index": i, "relevance_score": float(scores[i]),
                                                       "document": {"text": docs[i]}} for i in order],
                "usage": {"total_tokens": sum(map(len, ids))}}

PoolingService accepts text, IDs or a bounded batch of either. Its lock permits one device forward at a time; asyncio.to_thread keeps that forward off the HTTP event loop. The maximum batch and context limits bound each allocation. This is a separate padded encoder path, not continuous batching across HTTP embedding requests.

Chapter 36 already had an /v1/embeddings hook. Wire it to a trained local encoder:

from transformers import AutoModel
from izh.pooling import HFEncoder, PoolingService

# `server` is Chapter 36's Server; `tokenizer` must implement this model's text recipe.
encoder = AutoModel.from_pretrained("models/my-embedding-model", local_files_only=True).eval()
server.embedder = PoolingService(HFEncoder(encoder), tokenizer, server.model_name,
                                mode="mean", normalize=True, max_model_len=512)

The response preserves input order, counts non-padding prompt tokens and supports float lists or base64 little-endian FP32 vectors. Arbitrary dimension truncation is refused: it is a training property, not a free compression operation.

Reranking pairs a query with every candidate document, processes each pair jointly, then applies a trained scalar head to a pooled hidden vector. Unlike independent embeddings, query and document tokens interact inside attention. It costs one forward per pair and usually follows a cheap embedding search over many documents. Set server.reranker to a PoolingService with score_head, and provide tokenizer.encode_pair(query, document) with the trained pair convention. /v1/rerank returns scores sorted by descending relevance and original indices; it is an extension, not an OpenAI standard endpoint. A score need not be a calibrated probability.

Many adapters, one base-model batch

For token row $i$ selecting adapter $s_i$, Chapter 22’s equation becomes

$$y_i = W x_i + \frac{\alpha_{s_i}}{r_{s_i}} B_{s_i}(A_{s_i}x_i).$$

Compute the base projection over all rows once. Then shrink each row from width $D$ to its adapter’s rank, expand to the output width and add. Rows selecting no adapter get zero correction. Padding adapters to a common maximum rank lets the device buffers have fixed shapes; unused rows of A and columns of B are zero. The scale is folded into B at load time.

def grouped_lora(x, a, b, slots):
    """x [N,K], A [S,R,K], B [S,D,R], slots [N]; -1 means base only.
    Group rows by adapter, shrink then expand, and scatter to original order.  (Your engine: Chapter 43)
    B already includes alpha/r. Padded ranks have zero weights.
    """
    out = x.new_zeros((len(x), b.shape[1]))
    for slot in range(a.shape[0]):
        rows = torch.where(slots == slot)[0]
        if rows.numel():
            out[rows] = F.linear(F.linear(x[rows], a[slot]), b[slot])
    return out

The PyTorch reference groups rows by adapter and scatters them back. Two kernel shapes cover serving:

  • BGMV, batched gathered matrix-vector multiplication, handles decode rows, each with its own adapter slot.
  • SGMV, segmented gathered matrix-vector multiplication, sorts prefill rows by adapter so several rows reuse the same weights. Segment offsets describe the runs.
@triton.jit
def bgmv_kernel(X, W, IDS, Y, K: tl.constexpr, D: tl.constexpr, BK: tl.constexpr):
    """One program per row and output coordinate, selected adapter's matrix.  (Your engine: Chapter 43)"""
    row, out = tl.program_id(0), tl.program_id(1)
    slot = tl.load(IDS + row)
    k = tl.arange(0, BK)
    acc = tl.full((), 0, tl.float32)
    for start in range(0, K, BK):
        kk = start + k
        x = tl.load(X + row * K + kk, kk < K, other=0).to(tl.float32)
        w = tl.load(W + (tl.maximum(slot, 0) * D + out) * K + kk,
                    (kk < K) & (slot >= 0), other=0).to(tl.float32)
        acc += tl.sum(x * w, 0)
    tl.store(Y + row * D + out, acc)
@triton.jit
def sgmv_kernel(X, W, STARTS, Y, K: tl.constexpr, D: tl.constexpr, BK: tl.constexpr, BR: tl.constexpr):
    """Each adapter owns a contiguous segment; process BR rows together.  (Your engine: Chapter 43)"""
    tile, out, slot = tl.program_id(0), tl.program_id(1), tl.program_id(2)
    start, end = tl.load(STARTS + slot), tl.load(STARTS + slot + 1)
    rows = start + tile * BR + tl.arange(0, BR)
    k = tl.arange(0, BK)
    acc = tl.zeros((BR,), tl.float32)
    for base in range(0, K, BK):
        kk = base + k
        x = tl.load(X + rows[:, None] * K + kk[None, :],
                    (rows[:, None] < end) & (kk[None, :] < K), other=0).to(tl.float32)
        w = tl.load(W + (slot * D + out) * K + kk, kk < K, other=0).to(tl.float32)
        acc += tl.sum(x * w[None, :], 1)
    tl.store(Y + rows * D + out, acc, rows < end)

The kernels accumulate in FP32, mask odd widths and treat slot -1 as base-only. They perform both shrink and expand. These are readable vector-reduction kernels; the SGMV wrapper sorts at each projection and launches empty segments. Production implementations reuse grouping metadata and use tiled tensor-core GEMMs. Punica develops this batching approach; S-LoRA adds unified paging and adapter management at much larger scale.

Slots are storage; identities are content

class AdapterBank:
    """A fixed device budget for adapters; slots cannot be overwritten while leased."""
    def __init__(self, flat, capacity=8, max_rank=16, targets=("q_proj", "v_proj"), backend="torch"):
        if capacity < 1 or max_rank < 1 or backend not in ("torch", "triton"):
            raise ValueError("Need positive capacity/rank and backend torch|triton")
        if getattr(flat, "adapters", None) is not None:
            raise ValueError("An adapter bank is already installed")
        self.capacity, self.max_rank, self.backend = capacity, max_rank, backend
        self.names, self.keys, self.refs, self.layers = {}, {}, [0] * capacity, {}
        self.rows, self.decode = None, False
        for name, module in list(flat.model.named_modules()):
            if isinstance(module, nn.Linear) and name.rsplit(".", 1)[-1] in targets:
                parent, _, attr = name.rpartition(".")
                wrapped = AdapterLinear(module, self)
                setattr(flat.model.get_submodule(parent), attr, wrapped)
                self.layers[name] = wrapped
        if not self.layers:
            raise ValueError("No adapter targets found; install after projection fusion with matching target names")
        flat.adapters = self

    @torch.no_grad()
    def load(self, name, weights, alpha=None):
        """weights: {module_name: (A [rank,in], B [out,rank])}. Validate all before writing.
        A new generation's content digest prevents stale prefix-cache hits.  (Your engine: Chapter 43)
        """
        if not name or not weights or set(weights) - set(self.layers):
            raise ValueError("An adapter needs a name and known target modules")
        slot = self.names.get(name)
        if slot is not None and self.refs[slot]:
            raise ValueError("Cannot replace an adapter used by active requests")
        if slot is None:
            used = set(self.names.values())
            slot = next((s for s in range(self.capacity) if s not in used), None)
            if slot is None:
                raise ValueError("Adapter slots full; unload an idle adapter first")
        digest = hashlib.sha256()
        prepared = {}
        for target, (a, b) in sorted(weights.items()):
            layer = self.layers[target]
            rank = a.shape[0]
            if rank < 1 or rank > self.max_rank or a.shape != (rank, layer.in_features) \
                    or b.shape != (layer.out_features, rank):
                raise ValueError(f"{target}: incompatible LoRA shapes or rank")
            scale = (alpha if alpha is not None else rank) / rank
            a, b = a.detach().to(layer.a), b.detach().to(layer.b) * scale
            if not torch.isfinite(a).all() or not torch.isfinite(b).all():
                raise ValueError("Adapter contains non-finite weights")
            for tensor in (a, b):
                digest.update(str((target, tensor.shape, tensor.dtype)).encode())
                digest.update(tensor.contiguous().cpu().view(torch.uint8).numpy().tobytes())
            prepared[target] = a, b
        for target, layer in self.layers.items():
            layer.a[slot].zero_(), layer.b[slot].zero_()
            if target in prepared:
                a, b = prepared[target]
                layer.a[slot, :len(a)].copy_(a)
                layer.b[slot, :, :len(a)].copy_(b)
        self.names[name], self.keys[name] = slot, digest.hexdigest()
        return slot

    def load_peft(self, name, directory):
        """The Chapter 22 PEFT safetensors layout, with explicit rejection of extra features."""
        from ..safetensors_io import load_file
        path = Path(directory)
        config = json.loads((path / "adapter_config.json").read_text())
        if config.get("peft_type") != "LORA" or config.get("use_dora") or config.get("use_rslora") \
                or config.get("bias", "none") != "none" or config.get("modules_to_save") \
                or config.get("rank_pattern") or config.get("alpha_pattern"):
            raise ValueError("Only plain LoRA with uniform alpha, no trained bias or saved modules is supported")
        tensors = load_file(path / "adapter_model.safetensors")
        weights = {}
        for key in tensors:
            if not key.endswith(".lora_A.weight"):
                continue
            target = key.removeprefix("base_model.model.").removesuffix(".lora_A.weight")
            weights[target] = tensors[key], tensors[key.replace(".lora_A.", ".lora_B.")]
        expected = {f"base_model.model.{t}.lora_{part}.weight" for t in weights for part in ("A", "B")}
        if set(tensors) != expected:
            raise ValueError("Unsupported or incomplete adapter tensors")
        return self.load(name, weights, config["lora_alpha"])

    def unload(self, name):
        slot = self.names[name]
        if self.refs[slot]:
            raise ValueError("Cannot unload an adapter used by active requests")
        del self.names[name], self.keys[name]

    def acquire(self, name):
        if name not in self.names:
            raise ValueError(f"Unknown LoRA adapter {name!r}")
        slot = self.names[name]
        self.refs[slot] += 1
        return slot, ("lora", self.keys[name])

    def release(self, slot):
        if self.refs[slot] < 1:
            raise RuntimeError("Adapter lease released twice")
        self.refs[slot] -= 1

    @contextmanager
    def activate(self, rows, decode=False):
        previous = self.rows, self.decode
        self.rows, self.decode = rows, decode
        try:
            yield
        finally:
            self.rows, self.decode = previous

AdapterBank wraps selected linears with resident slot buffers. Loading validates every target and shape before writing any weights. Requests lease a slot at admission and release it on completion or cancellation; a preempted request keeps its lease. Replacing or unloading a leased adapter fails. Otherwise a request could prefill with adapter A and resume with B, keeping incompatible KV.

The prefix cache uses a digest of the adapter’s effective weights, not the integer slot. Reusing slot 0 for a different adapter must not reuse slot 0’s old KV. Conversely, two adapter names with identical effective weights can safely share prefixes. Changing base-model weights still requires draining the engine and clearing its caches.

load_peft reads the plain LoRA safetensors format used in Chapter 22. It refuses DoRA, rank-stabilized LoRA, trained biases, saved full modules and per-layer rank/alpha patterns. The launcher preloads named adapters:

python serve.py --model-dir models/Qwen3-0.6B --device cpu --dtype fp32 \
  --lora support=runs/support-adapter --lora coding=runs/coding-adapter

Requests use a name: {"prompt":"...", "lora":"support", "max_tokens":32}. The name and feature inputs travel through the process queue; the core derives the slot and cache salt. Clients cannot supply slot numbers or load arbitrary paths. Each API choice from n>1 gets its own lease, while sharing compatible prefix blocks.

Install the bank after any projection fusion, targeting the resulting module names. The launcher uses unfused Q/V targets. Quantized adapter bases, expert adapters and distributed adapter sharding need further work. Feature batches take the eager runner path because graph buffers don’t yet hold these inputs; feature requests with a drafter are refused until the draft path also consumes them. The dense paged cache, batching, prefix reuse and preemption remain shared with text generation.

Build it

Engine milestone 43: beyond text. Implement:

  • in engine/multimodal.py: VisionEncoder.forward, expand_images and mrope_cos_sin;
  • in engine/serve/runner.py: ModelRunner.prepare_features;
  • in engine/serve/adapters.py: grouped_lora and AdapterBank.load;
  • in engine/kernels/triton_lora.py: bgmv_kernel and sgmv_kernel;
  • in engine/pooling.py: pool_hidden.

Then run:

pytest tests/test_ch43_multimodal_lora_embeddings.py
python run.py features

The tests check ViT batch equivalence; exact text RoPE equivalence and axis selection; image generation against a separate cache-free forward; image chunks mixed with text; warm-image prefix hits and changed-image misses; mixed adapters against separately merged models in synchronous and asynchronous scheduling; slot leases and cancellation; odd-shaped shrink/expand kernels; padding-independent embeddings, float/base64 HTTP responses and reranking order.

Stretch exercises

  1. ★★ Add a trained VLM’s patch merger and processor. Compare image rows, rotary positions and logits with its official implementation before testing text answers. Where: VisionEncoder / VisionProcessor in engine/multimodal.py, with feature insertion in engine/serve/runner.py.
  2. ★★ Cache vision-encoder outputs by image bytes, preprocessing configuration and encoder revision. Bound the host cache; measure repeated-image latency separately from decoder prefix hits. Where: feature caching around VisionPrompt.features in engine/multimodal.py.
  3. ★★★ Replace the SGMV reductions with tiled matmuls and reuse the row permutation across layers. Measure crossover by adapter rank, tokens per adapter and number of adapters. Where: sgmv_kernel in engine/kernels/triton_lora.py, with permutation reuse in engine/serve/adapters.py.
  4. ★★ Add dynamic batching for embeddings with a bounded queue and a maximum hold time. Compare throughput and p99 latency against the current per-call lock. Where: PoolingService in engine/pooling.py, with endpoint dispatch in engine/serve/api.py.
  5. ★★★ Add adapter inputs to graph buffers, bucket by maximum rank, and test slot replacement after every request in a captured bucket retires. Where: adapter buffers in engine/serve/graphs.py and engine/serve/runner.py, with slot leases in engine/serve/adapters.py.

Check your understanding

  1. Why does a placeholder expand before the engine checks the context limit?
  2. Why do KV positions and M-RoPE positions differ after an image?
  3. What must an image prefix-cache key identify besides placeholder IDs?
  4. Why doesn’t mean pooling turn an ordinary decoder into a useful embedding model?
  5. Why must padding be excluded from both the pooled sum and its divisor?
  6. Which work is shared across requests with different LoRAs, and which work depends on the adapter?
  7. Why is an adapter slot number an unsafe prefix-cache identity?
  8. Why does a preempted request keep its adapter lease?

Going deeper

  • Dosovitskiy et al., An Image is Worth 16x16 Words, for ViT; Liu et al., Visual Instruction Tuning, for projecting image features into a language model.
  • Wang et al., Qwen2-VL, for M-RoPE and image/video position conventions.
  • Reimers and Gurevych, Sentence-BERT, for trained sentence embeddings and the distinction from cross-encoder scoring.
  • Chen et al., Punica; Sheng et al., S-LoRA, for gathered low-rank kernels and serving many adapters.