Skip to content

[PD] Charge each state component its own item length in the KV transfer metric - #33806

Open
AMD-yanfeiwang wants to merge 1 commit into
sgl-project:mainfrom
AMD-yanfeiwang:fix-pd-state-transfer-bytes
Open

AMD-yanfeiwang wants to merge 1 commit into
sgl-project:mainfrom
AMD-yanfeiwang:fix-pd-state-transfer-bytes

Conversation

@AMD-yanfeiwang

@AMD-yanfeiwang AMD-yanfeiwang commented Aug 6, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

CommonKVSender.get_transfer_metric() computes the state contribution as a product of two independently flattened sums:

total_bytes += self._transfer_num_state_indices * self.kv_mgr.state_item_lens_sum

KVArgs.state_item_lens is List[List[int]] — per component, then per layer — and it is flattened at construction:

self.state_item_lens_sum = sum(x for comp in args.state_item_lens for x in comp)

_record_transfer_indices flattens the index side the same way, adding len(component_indices) into a single counter without keeping track of which component it came from.

The product of those two is a cross product. With N components of index counts n_i and item-length sums L_i, the metric reports

(Σ n_i) · (Σ L_i)      instead of      Σ n_i · L_i

— inflated by exactly the cross terms Σ_{i≠j} n_i · L_j. Every index of every component is charged the item length of every other component as well.

It is correct only when a single component is in play, which is why test_kv_transfer_replica_metric.py (one component) never caught it.

Why it matters

The overcharge is worst when component sizes are lopsided, which is the normal case. A sliding-window ring ships a short index list; a compressed-state component ships a much longer one. Each ends up paying the other's price, and the two errors do not cancel.

It also makes the metric partly insensitive to transfer length: the cross terms scale with the other component's size rather than with what this component actually sent, so a large share of the reported bytes does not move when the real transfer does. That defeats the metric's main use — telling whether a change reduced KV traffic.

Concretely, with 2 components — 3 indices at 10 B and 1 index at 1000 B:

bytes
actual (3·10 + 1·1000) 1,030
reported ((3+1)·(10+1000)) 4,040

Change

Accumulate state bytes where the component association still exists, instead of trying to reconstruct them afterwards from two sums that have each lost it:

item_lens = self.kv_mgr.state_item_lens_per_component
for component_id, component_indices in enumerate(state_indices):
    if component_indices is None:
        continue
    self._transfer_state_bytes += len(component_indices) * item_lens[component_id]

The KV side keeps its scalar product — all layers ship the same page indices, so one KV index genuinely costs the sum over layers. Only the state side needed the per-component treatment.

state_item_lens_sum becomes state_item_lens_per_component, and the sender's _transfer_num_state_indices counter becomes _transfer_state_bytes. Both were private to this path; no other caller in the tree reads either.

Scope

Applies to every PD backend: mooncake, nixl and mori all call the inherited CommonKVSender._record_transfer_indices and none override get_transfer_metric. Notably the backends' data paths already index state_item_lens per component when computing offsets — only the accounting flattened it, so the transfers themselves were always correct.

Tests

test/registered/unit/disaggregation/test_kv_transfer_replica_metric.py:

  • test_each_state_component_is_charged_its_own_item_length — two components with deliberately lopsided item lengths; asserts the exact byte count and, separately, that it is not the flattened value.
  • test_a_component_with_no_indices_costs_nothing — a None slot must not inflate its neighbours' per-index price.

The existing replication cases are updated to the new field names and still pass unchanged in meaning.

Mutation-checked: restoring the flattened form fails exactly those two new cases and nothing else. Full test/registered/unit/disaggregation/ suite: 130 passed.

Interaction with #31993

#31993 (open) reworks this same function to fix a different problem — bytes being paired with a latency window that covers only the last chunk. It preserves the byte formula as-is, carrying the cross product into its new _transfer_bytes() helper:

def _transfer_bytes(self, num_kv_indices: int, num_state_indices: int) -> int:
    return (num_kv_indices * self.kv_item_lens_sum
            + num_state_indices * self.state_item_lens_sum)

The two changes are complementary — that one fixes when bytes are measured, this one fixes how many — but they touch adjacent lines and will conflict textually. Happy to rebase onto whichever lands first.


CI States

Latest PR Test (Base): ❌ Run #31069926741
Latest PR Test (Extra): ❌ Run #31069926654

`get_transfer_metric()` computed the state contribution as a product of two
independently flattened sums:

    self._transfer_num_state_indices * self.kv_mgr.state_item_lens_sum

`state_item_lens` is `List[List[int]]` -- per component, then per layer -- and
`_record_transfer_indices` collapsed the index side the same way, adding
`len(component_indices)` into one counter without keeping the component. The
product is therefore a cross product: every index of every component is charged
the item length of every OTHER component as well.

With N components of index counts n_i and item-length sums L_i, the reported
figure is (Σ n_i)(Σ L_i) instead of Σ n_i·L_i -- inflated by exactly the cross
terms. It is right only when a single component is in play, which is why the
existing test (one component) could not see it.

The overcharge is worst when component sizes are lopsided, which is the normal
case: a sliding-window ring ships a short index list against a much longer
compressed-state one, so each pays the other's price. It also makes the metric
partly insensitive to transfer length, because the cross terms scale with the
other component's size rather than with what this component actually sent.

Accumulate the bytes where the component is still known, rather than trying to
reconstruct them afterwards from two sums that have each lost the association.
The KV side keeps its scalar product: all layers ship the SAME page indices, so
one index really does cost the sum over layers.

Applies to every PD backend -- mooncake, nixl and mori all inherit
`CommonKVSender._record_transfer_indices` and none override it.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant