MCPcopy Create free account
hub / github.com/AlayaLab/Hive / _downsample_2d

Method _downsample_2d

models/flowsep/diffusers/models/resnet.py:360–412  ·  view source on GitHub ↗

Fused `Conv2d()` followed by `downsample_2d()`. Padding is performed only once at the beginning, not between the operations. The fused op is considerably more efficient than performing the same calculation using standard TensorFlow ops. It supports gradients of arbitrary o

(self, hidden_states, weight=None, kernel=None, factor=2, gain=1)

Source from the content-addressed store, hash-verified

358 self.out_channels = out_channels
359
360 def _downsample_2d(self, hidden_states, weight=None, kernel=None, factor=2, gain=1):
361 """Fused `Conv2d()` followed by `downsample_2d()`.
362 Padding is performed only once at the beginning, not between the operations. The fused op is considerably more
363 efficient than performing the same calculation using standard TensorFlow ops. It supports gradients of
364 arbitrary order.
365
366 Args:
367 hidden_states: Input tensor of the shape `[N, C, H, W]` or `[N, H, W, C]`.
368 weight:
369 Weight tensor of the shape `[filterH, filterW, inChannels, outChannels]`. Grouped convolution can be
370 performed by `inChannels = x.shape[0] // numGroups`.
371 kernel: FIR filter of the shape `[firH, firW]` or `[firN]` (separable). The default is `[1] *
372 factor`, which corresponds to average pooling.
373 factor: Integer downsampling factor (default: 2).
374 gain: Scaling factor for signal magnitude (default: 1.0).
375
376 Returns:
377 output: Tensor of the shape `[N, C, H // factor, W // factor]` or `[N, H // factor, W // factor, C]`, and
378 same datatype as `x`.
379 """
380
381 assert isinstance(factor, int) and factor >= 1
382 if kernel is None:
383 kernel = [1] * factor
384
385 # setup kernel
386 kernel = torch.tensor(kernel, dtype=torch.float32)
387 if kernel.ndim == 1:
388 kernel = torch.outer(kernel, kernel)
389 kernel /= torch.sum(kernel)
390
391 kernel = kernel * gain
392
393 if self.use_conv:
394 _, _, convH, convW = weight.shape
395 pad_value = (kernel.shape[0] - factor) + (convW - 1)
396 stride_value = [factor, factor]
397 upfirdn_input = upfirdn2d_native(
398 hidden_states,
399 torch.tensor(kernel, device=hidden_states.device),
400 pad=((pad_value + 1) // 2, pad_value // 2),
401 )
402 output = F.conv2d(upfirdn_input, weight, stride=stride_value, padding=0)
403 else:
404 pad_value = kernel.shape[0] - factor
405 output = upfirdn2d_native(
406 hidden_states,
407 torch.tensor(kernel, device=hidden_states.device),
408 down=factor,
409 pad=((pad_value + 1) // 2, pad_value // 2),
410 )
411
412 return output
413
414 def forward(self, hidden_states):
415 if self.use_conv:

Callers 1

forwardMethod · 0.95

Calls 1

upfirdn2d_nativeFunction · 0.85

Tested by

no test coverage detected