"""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