BLOCK_SIZE = 16 def block_hashes(token_ids, salt=None): """Chain-hash a token sequence into per-block keys.""" hashes, parent = [], hash(salt) # Only complete blocks are hashed. A partial tail block is skipped. for start in range(0, len(token_ids) - BLOCK_SIZE + 1, BLOCK_SIZE): block = tuple(token_ids[start : start + BLOCK_SIZE]) parent = hash((parent, block)) hashes.append(parent) return hashes def schedule(token_ids, cache): """Return how many tokens are reusable, and allocate the rest.""" matched_blocks = 0 for h in block_hashes(token_ids): if h not in cache: break # first miss ends all reuse cache[h].ref_count += 1 # pin it against eviction matched_blocks += 1 reused_tokens = matched_blocks * BLOCK_SIZE to_prefill = token_ids[reused_tokens:] return reused_tokens, to_prefill