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

Method build_mask

diffsynth/models/tiler.py:115–130  ·  view source on GitHub ↗
(self, data, is_bound)

Source from the content-addressed store, hash-verified

113
114
115 def build_mask(self, data, is_bound):
116 _, _, H, W = data.shape
117 h = repeat(torch.arange(H), "H -> H W", H=H, W=W)
118 w = repeat(torch.arange(W), "W -> H W", H=H, W=W)
119 border_width = (H + W) // 4
120 pad = torch.ones_like(h) * border_width
121 mask = torch.stack([
122 pad if is_bound[0] else h + 1,
123 pad if is_bound[1] else H - h,
124 pad if is_bound[2] else w + 1,
125 pad if is_bound[3] else W - w
126 ]).min(dim=0).values
127 mask = mask.clip(1, border_width)
128 mask = (mask / border_width).to(dtype=data.dtype, device=data.device)
129 mask = rearrange(mask, "H W -> 1 H W")
130 return mask
131
132
133 def tiled_forward(self, forward_fn, model_input, tile_size, tile_stride, tile_device="cpu", tile_dtype=torch.float32, border_width=None):

Callers 1

tiled_forwardMethod · 0.95

Calls 1

toMethod · 0.45

Tested by

no test coverage detected