Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -245,6 +245,7 @@ def start(self):
Headers.UNPAUSE,
Headers.SUSPEND,
Headers.RESUME,
Headers.INCREMENT_STALENESS,
Headers.STOP,
]:
# control signals for the engine
Expand Down
4 changes: 4 additions & 0 deletions megatron/core/inference/engines/dynamic_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -1617,6 +1617,10 @@ def schedule_requests(self) -> int:
self.suspend_signal = True
elif header == Headers.RESUME:
self.suspend_signal = False
elif header == Headers.INCREMENT_STALENESS:
waiting = set(self.waiting_request_ids)
for request_id, entry in self.requests.items():
entry.record.increment_staleness(policy_only=request_id in waiting)
elif header == Headers.STOP:
self.received_stop = True
else:
Expand Down
1 change: 1 addition & 0 deletions megatron/core/inference/headers.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@ class Headers(Enum):
UNPAUSE = auto()
SUSPEND = auto()
RESUME = auto()
INCREMENT_STALENESS = auto()
STOP = auto()
STOP_ACK = auto()

Expand Down
6 changes: 6 additions & 0 deletions megatron/core/inference/inference_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -213,13 +213,19 @@ def unpause_engines(self) -> None:
self.running.set()
self._send_signal_to_engines(Headers.UNPAUSE)

def increment_staleness(self):
"""Sends a signal to increment staleness on all in-flight requests."""
assert self.paused.is_set(), "Can only increment staleness while engines are paused."
self._send_signal_to_engines(Headers.INCREMENT_STALENESS)

def suspend_engines(self):
"""Sends a signal to pause all inference engines."""
self._send_signal_to_engines(Headers.PAUSE)
self._send_signal_to_engines(Headers.SUSPEND)

def resume_engines(self):
"""Sends a signal to unpause all inference engines."""
self.paused.clear()
self._send_signal_to_engines(Headers.RESUME)
self._send_signal_to_engines(Headers.UNPAUSE)

Expand Down
80 changes: 80 additions & 0 deletions megatron/core/inference/inference_request.py
Original file line number Diff line number Diff line change
Expand Up @@ -287,6 +287,8 @@ class DynamicInferenceRequest(InferenceRequest):
prompt_tokens: Optional[torch.Tensor] = None
# remaining prompt tokens are used for chunked prefill
remaining_prompt_tokens: Optional[torch.Tensor] = None
policy_staleness: Optional[torch.Tensor] = None
kv_cache_staleness: Optional[torch.Tensor] = None
latency: Optional[float] = None
# routing_indices stores MoE routing decisions for all tokens generated so far.
# Shape: [total_tokens, num_layers, topk] - accumulated across all generation steps
Expand Down Expand Up @@ -491,6 +493,53 @@ def request_id(self) -> int:
"""
return self.requests[0].request_id

@staticmethod
def _update_staleness_tensor(
tensor: Optional[torch.Tensor], total_tokens: int, increment: bool = True
) -> torch.Tensor:
"""Update a per-token staleness tensor, extending with zeros if needed.

Args:
tensor: Existing staleness tensor, or None to create a new one.
total_tokens: Expected length of the tensor after update.
increment: If True, increment all values by 1 (including new positions).
"""
if tensor is None:
tensor = torch.zeros(total_tokens, dtype=torch.int32, device='cpu')
elif len(tensor) < total_tokens:
tensor = torch.cat(
(
tensor,
torch.zeros(
total_tokens - len(tensor), dtype=tensor.dtype, device=tensor.device
),
),
dim=0,
)
if increment:
tensor = tensor + 1
return tensor

def increment_staleness(self, policy_only: bool = False):
"""Increment per-token staleness counters in-place.

Each call indicates that a training step has occurred since these tokens
were generated. Tokens not yet tracked are initialized to 1.

Args:
policy_only: If True, only increment policy_staleness. Use this for
evicted requests that have no KV cache to age.
"""
request = self[-1]
total_tokens = len(request.prompt_tokens) + len(request.generated_tokens)
request.policy_staleness = self._update_staleness_tensor(
request.policy_staleness, total_tokens, increment=True
)
if not policy_only:
request.kv_cache_staleness = self._update_staleness_tensor(
request.kv_cache_staleness, total_tokens, increment=True
)

def checkpoint(self, tokenizer: MegatronTokenizer | None = None):
"""Maintain reference to previous request, and then append a new request
that concatenates the previous prompt and generations.
Expand All @@ -501,6 +550,24 @@ def checkpoint(self, tokenizer: MegatronTokenizer | None = None):

old_request = self[-1]

total_tokens = len(old_request.prompt_tokens) + len(old_request.generated_tokens)

# Carry forward policy_staleness without incrementing.
policy_staleness = (
self._update_staleness_tensor(
old_request.policy_staleness, total_tokens, increment=False
)
if old_request.policy_staleness is not None
else None
)

# Reset kv_cache_staleness to 0.
kv_cache_staleness = (
self._update_staleness_tensor(None, total_tokens, increment=False)
if old_request.kv_cache_staleness is not None
else None
)

# New prompt (concatenate prompt + generated tokens).
new_prompt_tokens = torch.cat(
(
Expand Down Expand Up @@ -530,6 +597,8 @@ def checkpoint(self, tokenizer: MegatronTokenizer | None = None):
request_id=old_request.request_id,
prompt_tokens=new_prompt_tokens,
sampling_params=new_sampling_params,
policy_staleness=policy_staleness,
kv_cache_staleness=kv_cache_staleness,
)
# Preserve event_add_engine from old request if it exists, otherwise set it.
# This ensures TTFT calculation works correctly for evicted/resumed requests.
Expand Down Expand Up @@ -566,6 +635,15 @@ def merge_lists(key):
except TypeError as e: # generally means r.generated_text is None
generated_text = None

# Ensure staleness tensors are always materialized (zeros if never incremented).
total_tokens = len(prompt_tokens) + len(generated_tokens)
policy_staleness = self._update_staleness_tensor(
self.requests[-1].policy_staleness, total_tokens, increment=False
)
kv_cache_staleness = self._update_staleness_tensor(
self.requests[-1].kv_cache_staleness, total_tokens, increment=False
)

# Merged request.
request = DynamicInferenceRequest(
request_id=self.requests[0].request_id,
Expand All @@ -579,6 +657,8 @@ def merge_lists(key):
generated_log_probs=merge_lists("generated_log_probs"),
generated_top_n_logprobs=merge_lists("generated_top_n_logprobs"),
sampling_params=self.requests[0].sampling_params,
policy_staleness=policy_staleness,
kv_cache_staleness=kv_cache_staleness,
ttft=self.requests[0].ttft,
tpot=merge_lists("tpot"),
status=self.requests[-1].status,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -201,6 +201,15 @@ async def chat_completions():
"tool_calls" if metadata.get("tool_calls", []) else "stop"
), # Original code hardcoded this.
}
if result.get("policy_staleness") is not None:
choice_data["policy_staleness"] = result["policy_staleness"]
if result.get("kv_cache_staleness") is not None:
choice_data["kv_cache_staleness"] = result["kv_cache_staleness"]
events = result.get("events")
if events is not None:
num_evictions = sum(1 for e in events if e.get("type") == "EVICT")
if num_evictions > 0:
choice_data["num_evictions"] = num_evictions
if current_app.config['verbose']:
logging.info(result)
if result["routing_indices"] is not None:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -85,7 +85,7 @@ def health_check():
logger.info(f"Using parsers: {parsers}")

loop.set_default_executor(ThreadPoolExecutor(max_workers=8192))
await serve(AsyncioWSGIMiddleware(app), config)
await serve(AsyncioWSGIMiddleware(app, max_body_size=config.wsgi_max_body_size), config)


@trace_async_exceptions
Expand Down
137 changes: 137 additions & 0 deletions tests/unit_tests/inference/engines/test_dynamic_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -1847,3 +1847,140 @@ def test_suspend_resume_cycle(self, kv_cache_management_mode, static_kv_memory_p
f"Tensor address must be stable when static_kv_memory_pointers is set. "
f"Before: {addr_before:#x}, After: {addr_after:#x}"
)

@pytest.mark.skipif(
not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching"
)
@pytest.mark.parametrize("use_checkpoint", [False, True], ids=["persist", "recompute"])
@torch.inference_mode()
def test_staleness_tracking(self, use_checkpoint):
"""Test that staleness is correct tracked.
The use_checkpoint parameter simulates the behavior of different kv_cache_management_mode.
"""
PROMPT_LEN = 8
NUM_TOKENS = 8

test_config = DynamicEngineTestConfig(
num_requests=0,
min_prompt_length=PROMPT_LEN,
max_prompt_length=PROMPT_LEN,
num_tokens_to_generate=NUM_TOKENS,
)
env = self._build_test_env(test_config)
engine = env.engine

for i in range(2):
prompt_tokens = torch.randint(
0,
test_config.vocab_size - 1,
(PROMPT_LEN,),
dtype=torch.int64,
device=torch.cuda.current_device(),
)
engine._add_request(
DynamicInferenceRequest(
request_id=i,
prompt_tokens=prompt_tokens,
sampling_params=SamplingParams(
num_tokens_to_generate=NUM_TOKENS, termination_id=-1
),
)
)

for _ in range(3):
engine.step_modern()

for entry in engine.requests.values():
assert len(entry.record[-1].generated_tokens) == 3
assert entry.record[-1].policy_staleness is None
assert entry.record[-1].kv_cache_staleness is None

# Increment staleness.
for entry in engine.requests.values():
entry.record.increment_staleness()

for entry in engine.requests.values():
ps = entry.record[-1].policy_staleness
ks = entry.record[-1].kv_cache_staleness
assert ps.shape == ks.shape == (PROMPT_LEN + 3,)
assert (ps == 1).all()
assert (ks == 1).all()

# Simulate RECOMPUTE
if use_checkpoint:
for entry in engine.requests.values():
old_req = entry.record[-1]
event_add_engine = old_req.event_add_engine
entry.record.checkpoint()
# Prevent TTFT crash due to missing _add_request in test.
entry.record[-1].event_add_engine = event_add_engine

for entry in engine.requests.values():
ps = entry.record[-1].policy_staleness
ks = entry.record[-1].kv_cache_staleness
assert ps.shape == ks.shape == (PROMPT_LEN + 3,)
assert (ps == 1).all()
assert (ks == 0).all()

for _ in range(3):
engine.step_modern()

# Increment staleness.
for entry in engine.requests.values():
entry.record.increment_staleness()

for entry in engine.requests.values():
ps = entry.record[-1].policy_staleness
ks = entry.record[-1].kv_cache_staleness
assert ps.shape == ks.shape == (PROMPT_LEN + 6,)
assert (ps[: PROMPT_LEN + 3] == 2).all()
assert (ps[PROMPT_LEN + 3 :] == 1).all()
if use_checkpoint:
assert (ks == 1).all()
else:
assert (ks[: PROMPT_LEN + 3] == 2).all()
assert (ks[PROMPT_LEN + 3 :] == 1).all()

if use_checkpoint:
for entry in engine.requests.values():
old_req = entry.record[-1]
event_add_engine = old_req.event_add_engine
entry.record.checkpoint()
entry.record[-1].event_add_engine = event_add_engine

for entry in engine.requests.values():
ks = entry.record[-1].kv_cache_staleness
assert (ks == 0).all()

finished_records = []
while engine.has_unfinished_requests():
result = engine.step_modern()
finished_records.extend(result["finished_request_records"])

for record in finished_records:
merged = record.merge()

assert merged.policy_staleness is not None
assert merged.policy_staleness.shape == (PROMPT_LEN + NUM_TOKENS,)
assert (merged.policy_staleness[: PROMPT_LEN + 3] == 2).all()
assert (merged.policy_staleness[PROMPT_LEN + 3 : PROMPT_LEN + 6] == 1).all()
assert (merged.policy_staleness[PROMPT_LEN + 6 :] == 0).all()

assert merged.kv_cache_staleness is not None
assert merged.kv_cache_staleness.shape == (PROMPT_LEN + NUM_TOKENS,)
if use_checkpoint:
assert (merged.kv_cache_staleness == 0).all()
else:
assert (merged.kv_cache_staleness[: PROMPT_LEN + 3] == 2).all()
assert (merged.kv_cache_staleness[PROMPT_LEN + 3 : PROMPT_LEN + 6] == 1).all()
assert (merged.kv_cache_staleness[PROMPT_LEN + 6 :] == 0).all()

# Verify evicted requests don't have their policy staleness incremented.
record = finished_records[0]
record.checkpoint()
pre_ps = record[-1].policy_staleness.clone()

record.increment_staleness(policy_only=True) # This mimics the coordinator's action.

assert (record[-1].policy_staleness == pre_ps + 1).all()
assert (record[-1].kv_cache_staleness == 0).all()
Loading