refactor: centralize batch state and split scheduler duties
- make decode steady state self-validating: gate the token fill and sampling reuse on the decode cache's own task signature, drop TaskCacheManager.last_task_signature_matches whose req-index signature recycles with slot reuse - extract PolicyVersionGuard (version protocol plus generation/weight mutex) and Stepper (shared one-token advancement) out of the scheduler, keeping its public API unchanged - pin the input staging buffer only on CUDA devices so CPU-only workspaces allocate
This commit is contained in:
Vendored
-16
@@ -406,22 +406,6 @@ class TaskCacheManager:
|
||||
"""True if the last bind was a steady-state increment (same tasks, +1 seq_lens)."""
|
||||
return self._bind_was_steady
|
||||
|
||||
def last_task_signature_matches(self, task_ids: List[str]) -> bool:
|
||||
"""Check if task_ids match the previous bind's signature.
|
||||
|
||||
Used by Executor to detect steady-state decode for device-to-device
|
||||
token copy optimization.
|
||||
"""
|
||||
if self._bind_state is None:
|
||||
return False
|
||||
prev_sig = self._bind_state.sig
|
||||
# sig is tuple of req_indices, need to map task_ids to req_indices
|
||||
try:
|
||||
current_sig = tuple(self._states[tid].req_idx for tid in task_ids)
|
||||
return prev_sig == current_sig
|
||||
except KeyError:
|
||||
return False
|
||||
|
||||
# -- internals --
|
||||
|
||||
def _rollback(self, state: TaskCacheState, task_id: str):
|
||||
|
||||
@@ -415,14 +415,14 @@ class Executor:
|
||||
# inference-mode context.
|
||||
#
|
||||
# ``cache_valid`` checks the decode cache's own task signature:
|
||||
# req-index signatures in the cache manager are recycled when freed
|
||||
# slots are reallocated to new tasks, so a fresh batch whose prefill
|
||||
# re-bind coincides with a stale signature would otherwise replay a
|
||||
# previous generation's tokens into ``input_ids``.
|
||||
task_sig_match = self.task_cache.last_task_signature_matches(task_ids)
|
||||
# task ids are globally unique, so equality alone proves the cached
|
||||
# tokens were sampled for exactly this ordered batch. Req-index
|
||||
# signatures in the cache manager are deliberately NOT consulted —
|
||||
# they are recycled when freed slots are reallocated, which once
|
||||
# let a fresh batch replay a previous generation's tokens.
|
||||
cached = self._decode_cache
|
||||
cache_valid = cached is not None and cached.task_sig == task_sig
|
||||
if task_sig_match and cache_valid and cached.last_tokens is not None:
|
||||
if cache_valid and cached.last_tokens is not None:
|
||||
with torch.inference_mode():
|
||||
input_ids = ws.fill_input_ids_from_device(cached.last_tokens)
|
||||
else:
|
||||
@@ -433,12 +433,9 @@ class Executor:
|
||||
kv_cache = self.task_cache.bind(task_ids, ws)
|
||||
|
||||
# Reuse sampling state only if all conditions hold:
|
||||
# 1. KV bind detected steady increment (same req_indices, seq_lens +1)
|
||||
# 2. Task signature matches (same task_ids in same order)
|
||||
# 3. We have a valid cached decode state for THIS task set
|
||||
reuse_decode_state = (
|
||||
cache_valid and self.task_cache.bind_was_steady and task_sig_match
|
||||
)
|
||||
# 1. The cached decode state belongs to THIS task set (task_sig)
|
||||
# 2. KV bind detected steady increment (same req_indices, seq_lens +1)
|
||||
reuse_decode_state = cache_valid and self.task_cache.bind_was_steady
|
||||
if reuse_decode_state:
|
||||
info = cached.sampling_info
|
||||
ws.position_ids[:b] += 1
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
"""One-token advancement primitive shared by every scheduling mode."""
|
||||
|
||||
from contextlib import nullcontext
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
from astrai.extension import AttentionBackend, attn_backend
|
||||
from astrai.inference.cache import PagePool, TaskCacheManager
|
||||
from astrai.inference.metrics import MetricsCollector
|
||||
from astrai.inference.runtime.executor import Executor
|
||||
from astrai.inference.task import Task, TaskStatus
|
||||
|
||||
|
||||
class Stepper:
|
||||
"""Advance every active task by one token (prefill + decode).
|
||||
|
||||
Single shared primitive for both the continuous-batching loop and the
|
||||
synchronous ``run_batch`` path, so the two cannot drift.
|
||||
|
||||
Tasks must already be allocated in the KV cache. Tasks without output
|
||||
are prefilled first and sample their first token from the final prompt
|
||||
position. Tasks with output extend the cache by one position and decode
|
||||
from their latest generated token.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pool: PagePool,
|
||||
task_cache: TaskCacheManager,
|
||||
executor: Executor,
|
||||
metrics: MetricsCollector,
|
||||
):
|
||||
self._pool = pool
|
||||
self._task_cache = task_cache
|
||||
self._executor = executor
|
||||
self._metrics = metrics
|
||||
|
||||
@staticmethod
|
||||
def _task_backend_groups(tasks: List[Task]):
|
||||
groups = {}
|
||||
for task in tasks:
|
||||
groups.setdefault(task.backend, (task.backend, []))[1].append(task)
|
||||
return groups.values()
|
||||
|
||||
def step(
|
||||
self, tasks: List[Task], return_logprobs: bool = False
|
||||
) -> Tuple[List[Task], List[Task]]:
|
||||
"""Advance ``tasks`` by one token.
|
||||
|
||||
Args:
|
||||
tasks: Active tasks to advance by one token.
|
||||
return_logprobs: Forwarded to the executor; per-token logprobs
|
||||
are recorded on each task's ``output_logprobs``.
|
||||
|
||||
Returns:
|
||||
``(decoded, aborted)``: tasks that produced a new token (its ID
|
||||
already appended to ``output_ids``) and tasks that hit the
|
||||
sequence cap and were marked ``ABORTED``.
|
||||
"""
|
||||
to_prefill = [t for t in tasks if not t.prefill_done and t.prompt_ids]
|
||||
prefilled_ids = set()
|
||||
produced: List[Task] = []
|
||||
if to_prefill:
|
||||
for t in to_prefill:
|
||||
t.input_tokens = len(t.prompt_ids)
|
||||
|
||||
groups: Dict[Tuple[int, Optional[AttentionBackend]], List[Task]] = {}
|
||||
for t in to_prefill:
|
||||
start_pos = min(
|
||||
self._task_cache.task_cached(t.task_id), len(t.prompt_ids) - 1
|
||||
)
|
||||
groups.setdefault((start_pos, t.backend), []).append(t)
|
||||
|
||||
for (start_pos, _), group in groups.items():
|
||||
backend = group[0].backend
|
||||
backend_context = (
|
||||
attn_backend(backend) if backend is not None else nullcontext()
|
||||
)
|
||||
with (
|
||||
backend_context,
|
||||
self._metrics.record([t.task_id for t in group], "prefill"),
|
||||
):
|
||||
prefilled, step_out = self._executor.execute_prefill(
|
||||
group, start_pos=start_pos, return_logprobs=return_logprobs
|
||||
)
|
||||
|
||||
for t, out in zip(prefilled, step_out):
|
||||
t.output_ids.append(out[0] if return_logprobs else out)
|
||||
t.output_tokens += 1
|
||||
t.mark_prefill_done()
|
||||
prefilled_ids.add(t.task_id)
|
||||
produced.append(t)
|
||||
|
||||
start_logical_page = start_pos // self._pool.page_size
|
||||
for t in group:
|
||||
self._task_cache.task_record_hashes(
|
||||
t.task_id, t.prompt_ids, start_logical_page
|
||||
)
|
||||
|
||||
decoded: List[Task] = []
|
||||
aborted: List[Task] = []
|
||||
for t in tasks:
|
||||
if t.task_id in prefilled_ids:
|
||||
continue
|
||||
if self._task_cache.task_extend(t.task_id, t.next_pos):
|
||||
decoded.append(t)
|
||||
else:
|
||||
t.status = TaskStatus.ABORTED
|
||||
aborted.append(t)
|
||||
|
||||
for backend, group in self._task_backend_groups(decoded):
|
||||
backend_context = (
|
||||
attn_backend(backend) if backend is not None else nullcontext()
|
||||
)
|
||||
with (
|
||||
backend_context,
|
||||
self._metrics.record([t.task_id for t in group], "decode"),
|
||||
):
|
||||
step_out = self._executor.execute_decode(
|
||||
group, return_logprobs=return_logprobs
|
||||
)
|
||||
for t, out in zip(group, step_out):
|
||||
t.output_ids.append(out[0] if return_logprobs else out)
|
||||
t.output_tokens += 1
|
||||
t.advance_kv()
|
||||
produced.append(t)
|
||||
|
||||
return produced, aborted
|
||||
+25
-150
@@ -16,6 +16,7 @@ from astrai.extension import (
|
||||
from astrai.inference.cache import PagePool, TaskCacheManager
|
||||
from astrai.inference.metrics import MetricsCollector
|
||||
from astrai.inference.runtime.executor import Executor
|
||||
from astrai.inference.runtime.stepper import Stepper
|
||||
from astrai.inference.task import (
|
||||
STOP,
|
||||
GenerationResult,
|
||||
@@ -23,6 +24,7 @@ from astrai.inference.task import (
|
||||
TaskManager,
|
||||
TaskStatus,
|
||||
)
|
||||
from astrai.inference.versioning import PolicyVersionGuard
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||
|
||||
@@ -119,49 +121,32 @@ class InferenceScheduler:
|
||||
enable_cuda_graph=enable_cuda_graph,
|
||||
)
|
||||
|
||||
self._stepper = Stepper(
|
||||
self._cache, self._task_cache, self._executor, self._metrics
|
||||
)
|
||||
|
||||
self._stop_event = threading.Event()
|
||||
self._loop_thread: Optional[threading.Thread] = None
|
||||
self._weight_lock = threading.RLock()
|
||||
self._policy_version = policy_version
|
||||
self._policy_guard = PolicyVersionGuard(
|
||||
policy_version,
|
||||
ensure_ready=self._ensure_weight_update_ready,
|
||||
on_commit=self._task_cache.invalidate_cache,
|
||||
)
|
||||
# Synchronous generation shares the guard's generation/weight mutex.
|
||||
self._weight_lock = self._policy_guard.lock
|
||||
|
||||
@property
|
||||
def policy_version(self) -> int:
|
||||
"""Version of the model weights used for subsequent generations."""
|
||||
return self._policy_version
|
||||
|
||||
def _validate_weight_version(
|
||||
self, policy_version: int, *, require_advance: bool = False
|
||||
) -> None:
|
||||
if (
|
||||
isinstance(policy_version, bool)
|
||||
or not isinstance(policy_version, int)
|
||||
or policy_version < 0
|
||||
):
|
||||
raise ValueError("policy_version must be a non-negative integer")
|
||||
if policy_version < self._policy_version:
|
||||
raise ValueError(
|
||||
f"policy_version cannot move backwards from "
|
||||
f"{self._policy_version} to {policy_version}"
|
||||
)
|
||||
if require_advance and policy_version == self._policy_version:
|
||||
raise ValueError(
|
||||
f"policy_version must advance beyond {self._policy_version} "
|
||||
"when model weights are mutated"
|
||||
)
|
||||
return self._policy_guard.policy_version
|
||||
|
||||
def _ensure_weight_update_ready(self) -> None:
|
||||
"""Check weight update preconditions. Must be called under _weight_lock."""
|
||||
"""Check weight update preconditions. Must be called under the lock."""
|
||||
if self._loop_thread is not None and self._loop_thread.is_alive():
|
||||
raise RuntimeError("Stop the scheduler before updating model weights")
|
||||
if self._task_mgr.get_active_tasks() or self._task_mgr.get_waiting_tasks():
|
||||
raise RuntimeError("Cannot update model weights while tasks are queued")
|
||||
|
||||
def _commit_weight_version(self, policy_version: int) -> int:
|
||||
self._task_cache.invalidate_cache()
|
||||
self._policy_version = policy_version
|
||||
return self._policy_version
|
||||
|
||||
@_with_weight_lock
|
||||
def update_weights(self, policy_version: int) -> int:
|
||||
"""Acknowledge an in-place weight update and invalidate stale KV state.
|
||||
|
||||
@@ -170,13 +155,8 @@ class InferenceScheduler:
|
||||
version update makes that lifecycle visible and prevents prefix KV
|
||||
entries produced by older weights from being reused.
|
||||
"""
|
||||
self._validate_weight_version(policy_version)
|
||||
if policy_version == self._policy_version:
|
||||
return self._policy_version
|
||||
self._ensure_weight_update_ready()
|
||||
return self._commit_weight_version(policy_version)
|
||||
return self._policy_guard.update_weights(policy_version)
|
||||
|
||||
@_with_weight_lock
|
||||
def apply_weight_update(
|
||||
self, policy_version: Optional[int], update: Callable[[], T]
|
||||
) -> T:
|
||||
@@ -186,24 +166,11 @@ class InferenceScheduler:
|
||||
callers that only need "advance by one" (e.g. ``optimizer.step()``)
|
||||
without a read-compute-write race on the current version.
|
||||
"""
|
||||
if not callable(update):
|
||||
raise TypeError("update must be callable")
|
||||
if policy_version is None:
|
||||
policy_version = self._policy_version + 1
|
||||
else:
|
||||
self._validate_weight_version(policy_version, require_advance=True)
|
||||
self._ensure_weight_update_ready()
|
||||
return self._policy_guard.apply_weight_update(policy_version, update)
|
||||
|
||||
result = update()
|
||||
self._commit_weight_version(policy_version)
|
||||
return result
|
||||
|
||||
@_with_weight_lock
|
||||
def with_policy_snapshot(self, inspect: Callable[[int], T]) -> T:
|
||||
"""Inspect state while the scheduler's policy version remains stable."""
|
||||
if not callable(inspect):
|
||||
raise TypeError("inspect must be callable")
|
||||
return inspect(self._policy_version)
|
||||
return self._policy_guard.with_policy_snapshot(inspect)
|
||||
|
||||
def add_task(self, prompt: str, **kwargs) -> str:
|
||||
return self._task_mgr.add_task(prompt, **kwargs)
|
||||
@@ -226,7 +193,7 @@ class InferenceScheduler:
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
stats = self._task_mgr.get_stats()
|
||||
stats["kv_cache_tasks"] = self._task_cache.task_count
|
||||
stats["policy_version"] = self._policy_version
|
||||
stats["policy_version"] = self._policy_guard.policy_version
|
||||
return stats
|
||||
|
||||
@property
|
||||
@@ -242,105 +209,11 @@ class InferenceScheduler:
|
||||
return nullcontext()
|
||||
return attn_backend(self._backend)
|
||||
|
||||
@staticmethod
|
||||
def _task_backend_groups(tasks: List[Task]):
|
||||
groups = {}
|
||||
for task in tasks:
|
||||
groups.setdefault(task.backend, (task.backend, []))[1].append(task)
|
||||
return groups.values()
|
||||
|
||||
def _step(
|
||||
self, tasks: List[Task], return_logprobs: bool = False
|
||||
) -> Tuple[List[Task], List[Task]]:
|
||||
"""Advance every active task by one token (prefill + decode).
|
||||
|
||||
Single shared primitive for both the continuous-batching loop and
|
||||
the synchronous ``run_batch`` path, so the two cannot drift.
|
||||
|
||||
Tasks must already be allocated in the KV cache. Tasks without output
|
||||
are prefilled first and sample their first token from the final prompt
|
||||
position. Tasks with output extend the cache by one position and decode
|
||||
from their latest generated token.
|
||||
|
||||
Args:
|
||||
tasks: Active tasks to advance by one token.
|
||||
return_logprobs: Forwarded to ``execute_decode``; per-token
|
||||
logprobs are recorded on each task's ``output_logprobs``.
|
||||
|
||||
Returns:
|
||||
``(decoded, aborted)``: tasks that produced a new token (its ID
|
||||
already appended to ``output_ids``) and tasks that hit the
|
||||
sequence cap and were marked ``ABORTED``.
|
||||
"""
|
||||
to_prefill = [t for t in tasks if not t.prefill_done and t.prompt_ids]
|
||||
prefilled_ids = set()
|
||||
produced: List[Task] = []
|
||||
if to_prefill:
|
||||
for t in to_prefill:
|
||||
t.input_tokens = len(t.prompt_ids)
|
||||
|
||||
groups: Dict[Tuple[int, Optional[AttentionBackend]], List[Task]] = {}
|
||||
for t in to_prefill:
|
||||
start_pos = min(
|
||||
self._task_cache.task_cached(t.task_id), len(t.prompt_ids) - 1
|
||||
)
|
||||
groups.setdefault((start_pos, t.backend), []).append(t)
|
||||
|
||||
for (start_pos, _), group in groups.items():
|
||||
backend = group[0].backend
|
||||
backend_context = (
|
||||
attn_backend(backend) if backend is not None else nullcontext()
|
||||
)
|
||||
with (
|
||||
backend_context,
|
||||
self._metrics.record([t.task_id for t in group], "prefill"),
|
||||
):
|
||||
prefilled, step_out = self._executor.execute_prefill(
|
||||
group, start_pos=start_pos, return_logprobs=return_logprobs
|
||||
)
|
||||
|
||||
for t, out in zip(prefilled, step_out):
|
||||
t.output_ids.append(out[0] if return_logprobs else out)
|
||||
t.output_tokens += 1
|
||||
t.mark_prefill_done()
|
||||
prefilled_ids.add(t.task_id)
|
||||
produced.append(t)
|
||||
|
||||
start_logical_page = start_pos // self._cache.page_size
|
||||
for t in group:
|
||||
self._task_cache.task_record_hashes(
|
||||
t.task_id, t.prompt_ids, start_logical_page
|
||||
)
|
||||
|
||||
decoded: List[Task] = []
|
||||
aborted: List[Task] = []
|
||||
for t in tasks:
|
||||
if t.task_id in prefilled_ids:
|
||||
continue
|
||||
if self._task_cache.task_extend(t.task_id, t.next_pos):
|
||||
decoded.append(t)
|
||||
else:
|
||||
t.status = TaskStatus.ABORTED
|
||||
aborted.append(t)
|
||||
|
||||
for backend, group in self._task_backend_groups(decoded):
|
||||
backend_context = (
|
||||
attn_backend(backend) if backend is not None else nullcontext()
|
||||
)
|
||||
with (
|
||||
backend_context,
|
||||
self._metrics.record([t.task_id for t in group], "decode"),
|
||||
):
|
||||
step_out = self._executor.execute_decode(
|
||||
group, return_logprobs=return_logprobs
|
||||
)
|
||||
for t, out in zip(group, step_out):
|
||||
t.output_ids.append(out[0] if return_logprobs else out)
|
||||
t.output_tokens += 1
|
||||
t.advance_kv()
|
||||
produced.append(t)
|
||||
|
||||
return produced, aborted
|
||||
"""Advance every active task by one token; see :class:`Stepper`."""
|
||||
return self._stepper.step(tasks, return_logprobs=return_logprobs)
|
||||
|
||||
def _run_generation_loop(self):
|
||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||
@@ -389,7 +262,7 @@ class InferenceScheduler:
|
||||
if task.status != TaskStatus.ABORTED
|
||||
]
|
||||
|
||||
decoded, aborted = self._step(active)
|
||||
decoded, aborted = self._stepper.step(active)
|
||||
|
||||
for t in aborted:
|
||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||
@@ -533,7 +406,9 @@ class InferenceScheduler:
|
||||
|
||||
with self._backend_context():
|
||||
while live:
|
||||
decoded, aborted = self._step(live, return_logprobs=return_logprobs)
|
||||
decoded, aborted = self._stepper.step(
|
||||
live, return_logprobs=return_logprobs
|
||||
)
|
||||
for task in aborted:
|
||||
runtime_errors[task.task_id] = "kv_cache_extension_failed"
|
||||
live = [t for t in decoded if not t.is_finished(stop_ids)]
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
"""Policy version protocol for weights shared between trainer and server.
|
||||
|
||||
The guard owns the monotonic version counter, the RLock that serializes
|
||||
weight publication against generation (``run_batch`` acquires the same
|
||||
lock), and the validation/commit rules around both. Scheduler-specific
|
||||
preconditions (no in-flight generation, no queued tasks) and side effects
|
||||
(dropping stale KV entries) are injected as callables so the guard stays
|
||||
free of inference-subsystem knowledge.
|
||||
"""
|
||||
|
||||
import threading
|
||||
from functools import wraps
|
||||
from typing import Callable, Optional, TypeVar
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def _locked(method):
|
||||
@wraps(method)
|
||||
def synchronized(self, *args, **kwargs):
|
||||
with self._lock:
|
||||
return method(self, *args, **kwargs)
|
||||
|
||||
return synchronized
|
||||
|
||||
|
||||
class PolicyVersionGuard:
|
||||
"""Monotonic policy-version protocol over shared in-place weights.
|
||||
|
||||
The scheduler and the in-process trainer mutate the same model object;
|
||||
every weight mutation must publish a new version atomically under the
|
||||
generation lock. Versions never move backwards and every mutation
|
||||
advances (``apply_weight_update``) or repeats the live version
|
||||
(``update_weights`` idempotently).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
policy_version: int,
|
||||
ensure_ready: Callable[[], None],
|
||||
on_commit: Callable[[], None],
|
||||
):
|
||||
if (
|
||||
isinstance(policy_version, bool)
|
||||
or not isinstance(policy_version, int)
|
||||
or policy_version < 0
|
||||
):
|
||||
raise ValueError("policy_version must be a non-negative integer")
|
||||
self._lock = threading.RLock()
|
||||
self._policy_version = policy_version
|
||||
self._ensure_ready = ensure_ready
|
||||
self._on_commit = on_commit
|
||||
|
||||
@property
|
||||
def lock(self) -> threading.RLock:
|
||||
"""Generation/weight mutex; synchronous generation acquires it too."""
|
||||
return self._lock
|
||||
|
||||
@property
|
||||
def policy_version(self) -> int:
|
||||
"""Version of the model weights used for subsequent generations."""
|
||||
return self._policy_version
|
||||
|
||||
def _validate(self, policy_version: int, *, require_advance: bool = False) -> None:
|
||||
if (
|
||||
isinstance(policy_version, bool)
|
||||
or not isinstance(policy_version, int)
|
||||
or policy_version < 0
|
||||
):
|
||||
raise ValueError("policy_version must be a non-negative integer")
|
||||
if policy_version < self._policy_version:
|
||||
raise ValueError(
|
||||
f"policy_version cannot move backwards from "
|
||||
f"{self._policy_version} to {policy_version}"
|
||||
)
|
||||
if require_advance and policy_version == self._policy_version:
|
||||
raise ValueError(
|
||||
f"policy_version must advance beyond {self._policy_version} "
|
||||
"when model weights are mutated"
|
||||
)
|
||||
|
||||
@_locked
|
||||
def update_weights(self, policy_version: int) -> int:
|
||||
"""Acknowledge an in-place weight update and invalidate stale KV state.
|
||||
|
||||
The scheduler owns the same model object as the in-process trainer,
|
||||
so weights have already changed when this method is called. The
|
||||
explicit version update makes that lifecycle visible and prevents
|
||||
prefix KV entries produced by older weights from being reused.
|
||||
"""
|
||||
self._validate(policy_version)
|
||||
if policy_version == self._policy_version:
|
||||
return self._policy_version
|
||||
self._ensure_ready()
|
||||
return self._commit(policy_version)
|
||||
|
||||
@_locked
|
||||
def apply_weight_update(
|
||||
self, policy_version: Optional[int], update: Callable[[], T]
|
||||
) -> T:
|
||||
"""Mutate shared weights and publish their version without generation.
|
||||
|
||||
``policy_version=None`` derives ``live + 1`` under the same lock, for
|
||||
callers that only need "advance by one" (e.g. ``optimizer.step()``)
|
||||
without a read-compute-write race on the current version.
|
||||
"""
|
||||
if not callable(update):
|
||||
raise TypeError("update must be callable")
|
||||
if policy_version is None:
|
||||
policy_version = self._policy_version + 1
|
||||
else:
|
||||
self._validate(policy_version, require_advance=True)
|
||||
self._ensure_ready()
|
||||
|
||||
result = update()
|
||||
self._commit(policy_version)
|
||||
return result
|
||||
|
||||
@_locked
|
||||
def with_policy_snapshot(self, inspect: Callable[[int], T]) -> T:
|
||||
"""Inspect state while the policy version remains stable."""
|
||||
if not callable(inspect):
|
||||
raise TypeError("inspect must be callable")
|
||||
return inspect(self._policy_version)
|
||||
|
||||
def _commit(self, policy_version: int) -> int:
|
||||
self._on_commit()
|
||||
self._policy_version = policy_version
|
||||
return self._policy_version
|
||||
@@ -80,7 +80,10 @@ class InferenceWorkspace:
|
||||
(max_batch_size,), dtype=torch.long, device=device
|
||||
)
|
||||
self._pin = torch.empty(
|
||||
(max_batch_size,), dtype=torch.long, pin_memory=True
|
||||
(max_batch_size,),
|
||||
dtype=torch.long,
|
||||
pin_memory=torch.cuda.is_available()
|
||||
and torch.device(device).type == "cuda",
|
||||
)
|
||||
|
||||
# KV-cache bind metadata (fixed shape, written by
|
||||
|
||||
Reference in New Issue
Block a user