(self, hidden_states, timestep, prompt_emb, tile_size=(60, 90), tile_stride=(30, 45))
| 253 | |
| 254 | |
| 255 | def tiled_forward(self, hidden_states, timestep, prompt_emb, tile_size=(60, 90), tile_stride=(30, 45)): |
| 256 | B, C, T, H, W = hidden_states.shape |
| 257 | value = torch.zeros((B, C, T, H, W), dtype=hidden_states.dtype, device=hidden_states.device) |
| 258 | weight = torch.zeros((B, C, T, H, W), dtype=hidden_states.dtype, device=hidden_states.device) |
| 259 | |
| 260 | # Split tasks |
| 261 | tasks = [] |
| 262 | for h in range(0, H, tile_stride): |
| 263 | for w in range(0, W, tile_stride): |
| 264 | if (h-tile_stride >= 0 and h-tile_stride+tile_size >= H) or (w-tile_stride >= 0 and w-tile_stride+tile_size >= W): |
| 265 | continue |
| 266 | h_, w_ = h + tile_size, w + tile_size |
| 267 | if h_ > H: h, h_ = max(H - tile_size, 0), H |
| 268 | if w_ > W: w, w_ = max(W - tile_size, 0), W |
| 269 | tasks.append((h, h_, w, w_)) |
| 270 | |
| 271 | # Run |
| 272 | for hl, hr, wl, wr in tasks: |
| 273 | mask = self.build_mask( |
| 274 | value.shape[2], (hr-hl), (wr-wl), |
| 275 | hidden_states.dtype, hidden_states.device, |
| 276 | is_bound=(True, True, hl==0, hr>=H, wl==0, wr>=W) |
| 277 | ) |
| 278 | model_output = self.forward(hidden_states[:, :, :, hl:hr, wl:wr], timestep, prompt_emb) |
| 279 | value[:, :, :, hl:hr, wl:wr] += model_output * mask |
| 280 | weight[:, :, :, hl:hr, wl:wr] += mask |
| 281 | value = value / weight |
| 282 | |
| 283 | return value |
| 284 | |
| 285 | |
| 286 | def forward(self, hidden_states, timestep, prompt_emb, image_rotary_emb=None, tiled=False, tile_size=90, tile_stride=30, use_gradient_checkpointing=False): |
no test coverage detected