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