@@ -1288,7 +1288,14 @@ def __init__(
12881288 raise RuntimeError("failed to create MTP draft context")
12891289 ctx_other = llama_cpp_ext.llama_get_ctx_other(self.ctx)
12901290 self.is_mem_shared = bool(ctx_other and ctx_other == self.target_ctx)
1291- self.sampled_batch_draft = not self.is_mem_shared
1291+ self.n_mtp_layers = max(
1292+ 1,
1293+ int(llama_cpp.llama_model_n_layer_nextn(self.model)),
1294+ )
1295+ self.chain_heads = self.n_mtp_layers > 1 and not self.is_mem_shared
1296+ if self.chain_heads:
1297+ self.num_pred_tokens = min(self.num_pred_tokens, self.n_mtp_layers)
1298+ self.sampled_batch_draft = not self.is_mem_shared and not self.chain_heads
12921299 self.n_batch = int(llama_cpp.llama_n_batch(self.ctx))
12931300 mem = llama_cpp.llama_get_memory(self.ctx)
12941301 if mem is None:
@@ -1451,6 +1458,17 @@ def _try_decode_batch(self) -> bool:
14511458 return False
14521459 return True
14531460
1461+ def _set_nextn_layer_offset(self, offset: int) -> None:
1462+ if self.chain_heads:
1463+ llama_cpp_ext.llama_set_nextn_layer_offset(self.ctx, offset)
1464+
1465+ def _try_decode_batch_for_mtp_head(self, head: int) -> bool:
1466+ self._set_nextn_layer_offset(head)
1467+ try:
1468+ return self._try_decode_batch()
1469+ finally:
1470+ self._set_nextn_layer_offset(0)
1471+
14541472 def _decode_batch(self) -> None:
14551473 n_tokens = int(self.batch.n_tokens)
14561474 if n_tokens <= 0:
@@ -1464,6 +1482,22 @@ def _decode_batch(self) -> None:
14641482 self.decode_failures_total += 1
14651483 raise RuntimeError(f"MTP draft decode failed with code {result}")
14661484
1485+ def _decode_batch_for_mtp_heads(
1486+ self,
1487+ start_pos_by_seq: Dict[int, int],
1488+ ) -> None:
1489+ if not self.chain_heads:
1490+ self._decode_batch()
1491+ return
1492+ try:
1493+ for head in range(self.n_mtp_layers):
1494+ for seq_id, start_pos in start_pos_by_seq.items():
1495+ llama_cpp.llama_memory_seq_rm(self.mem, seq_id, start_pos, -1)
1496+ self._set_nextn_layer_offset(head)
1497+ self._decode_batch()
1498+ finally:
1499+ self._set_nextn_layer_offset(0)
1500+
14671501 def metric_definitions(
14681502 self,
14691503 ) -> List[Tuple[str, str, str, Union[int, float]]]:
@@ -1577,7 +1611,8 @@ def _process_rows(
15771611 target_rows_by_seq: Dict[int, List[int]],
15781612 aligned_by_seq: Dict[int, bool],
15791613 ) -> None:
1580- added_pos_by_seq: Dict[int, int] = {}
1614+ added_start_pos_by_seq: Dict[int, int] = {}
1615+ added_end_pos_by_seq: Dict[int, int] = {}
15811616 self._clear_batch()
15821617 for index in range(start, end):
15831618 if int(batch.n_seq_id[index]) != 1:
@@ -1609,15 +1644,85 @@ def _process_rows(
16091644 self._set_batch_embedding_row(slot, self.pending_h[seq_id])
16101645 else:
16111646 self._set_batch_embedding_row(slot, h_tgt_rows[previous_row])
1612- added_pos_by_seq[seq_id] = pos
1647+ added_start_pos_by_seq.setdefault(seq_id, pos)
1648+ added_end_pos_by_seq[seq_id] = pos
16131649 previous_row_by_seq[seq_id] = index
16141650 target_rows_by_seq.setdefault(seq_id, []).append(index)
16151651
16161652 if int(self.batch.n_tokens) > 0:
1617- self._decode_batch( )
1618- for seq_id, pos in added_pos_by_seq .items():
1653+ self._decode_batch_for_mtp_heads(added_start_pos_by_seq )
1654+ for seq_id, pos in added_end_pos_by_seq .items():
16191655 self.context_pos[seq_id] = max(self.context_pos[seq_id], pos + 1)
16201656
1657+ def _draft_chain_heads(
1658+ self,
1659+ *,
1660+ seq_id: int,
1661+ first_pos: int,
1662+ token: int,
1663+ n_predict: int,
1664+ ) -> np.ndarray:
1665+ if self.context_pos[seq_id] > first_pos:
1666+ self.truncate(seq_id, first_pos)
1667+ if self.context_pos[seq_id] < first_pos:
1668+ self.ready[seq_id] = False
1669+ return np.array([], dtype=np.intc)
1670+
1671+ drafted: List[int] = []
1672+ chain_tokens = [token]
1673+ chain_embeddings = [self.pending_h[seq_id].copy()]
1674+ self._reset_sampler(seq_id)
1675+
1676+ try:
1677+ for head in range(min(n_predict, self.n_mtp_layers)):
1678+ llama_cpp.llama_memory_seq_rm(self.mem, seq_id, first_pos, -1)
1679+ self._clear_batch()
1680+ for offset, (chain_token, embedding) in enumerate(
1681+ zip(chain_tokens, chain_embeddings)
1682+ ):
1683+ slot = int(self.batch.n_tokens)
1684+ self._add_batch_token(
1685+ token=chain_token,
1686+ pos=first_pos + offset,
1687+ seq_id=seq_id,
1688+ logits=offset == len(chain_tokens) - 1,
1689+ )
1690+ self._set_batch_embedding_row(slot, embedding)
1691+
1692+ output_index = int(self.batch.n_tokens) - 1
1693+ if not self._try_decode_batch_for_mtp_head(head):
1694+ break
1695+ self.context_pos[seq_id] = max(
1696+ self.context_pos[seq_id],
1697+ first_pos + len(chain_tokens),
1698+ )
1699+ sampled_token = self._sample_token(output_index, seq_id=seq_id)
1700+ if sampled_token is None:
1701+ break
1702+ drafted.append(sampled_token)
1703+ if len(drafted) >= n_predict:
1704+ break
1705+ h_row = llama_cpp_ext.llama_get_embeddings_nextn_ith(
1706+ self.ctx,
1707+ output_index,
1708+ )
1709+ if not h_row:
1710+ break
1711+ chain_tokens.append(sampled_token)
1712+ chain_embeddings.append(
1713+ np.ctypeslib.as_array(
1714+ h_row,
1715+ shape=(self.n_embd,),
1716+ ).copy()
1717+ )
1718+ finally:
1719+ self._set_nextn_layer_offset(0)
1720+ self.truncate(seq_id, first_pos)
1721+
1722+ if not drafted:
1723+ return np.array([], dtype=np.intc)
1724+ return np.asarray(drafted, dtype=np.intc)
1725+
16211726 def draft(
16221727 self,
16231728 input_ids: np.ndarray,
@@ -1649,6 +1754,13 @@ def draft(
16491754
16501755 token = int(input_ids[-1])
16511756 drafted: List[int] = []
1757+ if self.chain_heads:
1758+ return self._draft_chain_heads(
1759+ seq_id=seq_id,
1760+ first_pos=first_pos,
1761+ token=token,
1762+ n_predict=n_predict,
1763+ )
16521764 if not self.is_mem_shared and self.context_pos[seq_id] > first_pos:
16531765 self.truncate(seq_id, first_pos)
16541766 if not self.is_mem_shared and self.context_pos[seq_id] < first_pos:
@@ -1709,6 +1821,15 @@ def draft_many(
17091821 /,
17101822 ) -> List[np.ndarray]:
17111823 results = [np.array([], dtype=np.intc) for _ in requests]
1824+ if self.chain_heads:
1825+ for result_index, (input_ids, seq_id, max_tokens) in enumerate(requests):
1826+ results[result_index] = self.draft(
1827+ input_ids,
1828+ seq_id=seq_id,
1829+ max_tokens=max_tokens,
1830+ )
1831+ return results
1832+
17121833 active: List["MTPDraftProvider.DraftManyState"] = []
17131834 for result_index, (input_ids, seq_id, max_tokens) in enumerate(requests):
17141835 if (
0 commit comments