Skip to content

Commit a6a6dde

Browse files
committed
feat(example): support chained NextN server MTP
1 parent 4bee85b commit a6a6dde

2 files changed

Lines changed: 127 additions & 5 deletions

File tree

‎CHANGELOG.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
77

88
## [Unreleased]
99

10+
- feat(example): support chained NextN heads for server MTP drafting
1011
- feat: update llama.cpp to ggml-org/llama.cpp@92e854ab8
1112
- fix: preserve recurrent/hybrid model state when the full prompt is already cached by @allthatido and @abetlen in #2306
1213

‎examples/server/server.py‎

Lines changed: 126 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)