(
self, sample, timestep, encoder_hidden_states, add_time_id,
batch_time=25, batch_height=128, batch_width=128,
stride_time=5, stride_height=64, stride_width=64,
progress_bar=lambda x:x
)
| 335 | |
| 336 | |
| 337 | def tiled_forward( |
| 338 | self, sample, timestep, encoder_hidden_states, add_time_id, |
| 339 | batch_time=25, batch_height=128, batch_width=128, |
| 340 | stride_time=5, stride_height=64, stride_width=64, |
| 341 | progress_bar=lambda x:x |
| 342 | ): |
| 343 | data_device = sample.device |
| 344 | computation_device = self.conv_in.weight.device |
| 345 | torch_dtype = sample.dtype |
| 346 | T, C, H, W = sample.shape |
| 347 | |
| 348 | weight = torch.zeros((T, 1, H, W), dtype=torch_dtype, device=data_device) |
| 349 | values = torch.zeros((T, 4, H, W), dtype=torch_dtype, device=data_device) |
| 350 | |
| 351 | # Split tasks |
| 352 | tasks = [] |
| 353 | for t in range(0, T, stride_time): |
| 354 | for h in range(0, H, stride_height): |
| 355 | for w in range(0, W, stride_width): |
| 356 | if (t-stride_time >= 0 and t-stride_time+batch_time >= T)\ |
| 357 | or (h-stride_height >= 0 and h-stride_height+batch_height >= H)\ |
| 358 | or (w-stride_width >= 0 and w-stride_width+batch_width >= W): |
| 359 | continue |
| 360 | tasks.append((t, t+batch_time, h, h+batch_height, w, w+batch_width)) |
| 361 | |
| 362 | # Run |
| 363 | for tl, tr, hl, hr, wl, wr in progress_bar(tasks): |
| 364 | sample_batch = sample[tl:tr, :, hl:hr, wl:wr].to(computation_device) |
| 365 | sample_batch = self.forward(sample_batch, timestep, encoder_hidden_states, add_time_id).to(data_device) |
| 366 | mask = self.build_mask(sample_batch, is_bound=(tl==0, tr>=T, hl==0, hr>=H, wl==0, wr>=W)) |
| 367 | values[tl:tr, :, hl:hr, wl:wr] += sample_batch * mask |
| 368 | weight[tl:tr, :, hl:hr, wl:wr] += mask |
| 369 | values /= weight |
| 370 | return values |
| 371 | |
| 372 | |
| 373 | def forward(self, sample, timestep, encoder_hidden_states, add_time_id, use_gradient_checkpointing=False, **kwargs): |
nothing calls this directly
no test coverage detected