Nucleus-Resynthesis / runtime /src /resynthesis /bulk_wave_context.py
Wl6adams's picture
Add portable Release 188 generation runtime
919fd68 verified
Raw
History Blame Contribute Delete
987 Bytes
"""Per-wave bulk-training context for varlen ring and sequence-parallel canaries."""
from __future__ import annotations
import torch
_varlen_cu_seqlens: torch.Tensor | None = None
_varlen_max_seqlen: torch.Tensor | None = None
def set_bulk_varlen_context_boundary(
*,
cu_seqlens: torch.Tensor | None,
max_seqlen: torch.Tensor | None,
) -> None:
"""Install varlen metadata for the active bulk CUDA wave."""
global _varlen_cu_seqlens, _varlen_max_seqlen
_varlen_cu_seqlens = cu_seqlens
_varlen_max_seqlen = max_seqlen
def clear_bulk_varlen_context_boundary() -> None:
"""Drop varlen metadata after one wave completes."""
global _varlen_cu_seqlens, _varlen_max_seqlen
_varlen_cu_seqlens = None
_varlen_max_seqlen = None
def bulk_varlen_context_boundary() -> tuple[torch.Tensor | None, torch.Tensor | None]:
"""Return active varlen ``(cu_seqlens, max_seqlen)`` when installed."""
return _varlen_cu_seqlens, _varlen_max_seqlen