| 150 | return len(self.starts) - 1 |
| 151 | |
| 152 | def create_buffer(self): |
| 153 | distri_config = self.distri_config |
| 154 | if distri_config.rank == 0 and distri_config.verbose: |
| 155 | print( |
| 156 | f"Create buffer with {self.numel / 1e6:.3f}M parameters for {len(self.starts)} tensors on each device." |
| 157 | ) |
| 158 | for layer_type, numel in self.numel_dict.items(): |
| 159 | print(f" {layer_type}: {numel / 1e6:.3f}M parameters") |
| 160 | |
| 161 | self.buffer_list = [ |
| 162 | torch.empty(self.numel, dtype=self.torch_dtype, device=self.distri_config.device) |
| 163 | for _ in range(self.distri_config.n_device_per_batch) |
| 164 | ] |
| 165 | self.handles = [None for _ in range(len(self.starts))] |
| 166 | |
| 167 | def get_buffer_list(self, idx: int) -> List[torch.Tensor]: |
| 168 | buffer_list = [t[self.starts[idx] : self.ends[idx]].view(self.shapes[idx]) for t in self.buffer_list] |