MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / tiled_forward

Method tiled_forward

diffsynth/models/cog_dit.py:255–283  ·  view source on GitHub ↗
(self, hidden_states, timestep, prompt_emb, tile_size=(60, 90), tile_stride=(30, 45))

Source from the content-addressed store, hash-verified

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):

Callers 1

forwardMethod · 0.45

Calls 2

build_maskMethod · 0.95
forwardMethod · 0.95

Tested by

no test coverage detected