MCPcopy Create free account
hub / github.com/NVlabs/GSPN / create_buffer

Method create_buffer

t2i/src/distrifuser/utils.py:152–165  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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]

Callers 1

prepareMethod · 0.95

Calls 1

printFunction · 0.50

Tested by

no test coverage detected