Tensor parallel context.
| 12 | |
| 13 | @dataclass(frozen=True) |
| 14 | class TPInfo: |
| 15 | """Tensor parallel context.""" |
| 16 | |
| 17 | rank: int |
| 18 | size: int |
| 19 | |
| 20 | def is_rank0(self) -> bool: |
| 21 | return self.rank == 0 |
| 22 | |
| 23 | def rank0_print(self, *args, **kwargs) -> None: |
| 24 | if self.is_rank0(): |
| 25 | print(*args, **kwargs) |
| 26 | |
| 27 | @classmethod |
| 28 | def from_world(cls) -> "TPInfo": |
| 29 | return cls(rank=dist.get_rank(), size=dist.get_world_size()) |
| 30 | |
| 31 | |
| 32 | TP1 = TPInfo(rank=0, size=1) |
no outgoing calls
no test coverage detected