| import torch.nn.functional as F |
| import torch |
| import torch.nn as nn |
|
|
|
|
| def make_grid(input): |
| B, C, H, W = input.size() |
| |
| device = input.device |
| xx = torch.arange(0, W, device=device).view(1, -1).repeat(H, 1) |
| yy = torch.arange(0, H, device=device).view(-1, 1).repeat(1, W) |
| xx = xx.view(1, 1, H, W).repeat(B, 1, 1, 1) |
| yy = yy.view(1, 1, H, W).repeat(B, 1, 1, 1) |
| grid = torch.cat((xx, yy), 1).float() |
|
|
| return grid |
|
|
| def warp(input, flow, grid, mode="bilinear", padding_mode="zeros"): |
|
|
| B, C, H, W = input.size() |
| vgrid = grid + flow |
|
|
| vgrid[:, 0, :, :] = 2.0 * vgrid[:, 0, :, :].clone() / max(W - 1, 1) - 1.0 |
| vgrid[:, 1, :, :] = 2.0 * vgrid[:, 1, :, :].clone() / max(H - 1, 1) - 1.0 |
| vgrid = vgrid.permute(0, 2, 3, 1) |
| output = torch.nn.functional.grid_sample(input, vgrid, padding_mode=padding_mode, mode=mode, align_corners=True) |
| return output |
|
|
| def l2normalize(v, eps=1e-12): |
| return v / (v.norm() + eps) |
|
|
|
|
| class spectral_norm(nn.Module): |
| def __init__(self, module, name='weight', power_iterations=1): |
| super(spectral_norm, self).__init__() |
| self.module = module |
| self.name = name |
| self.power_iterations = power_iterations |
| if not self._made_params(): |
| self._make_params() |
|
|
| def _update_u_v(self): |
| u = getattr(self.module, self.name + "_u") |
| v = getattr(self.module, self.name + "_v") |
| w = getattr(self.module, self.name + "_bar") |
|
|
| height = w.data.shape[0] |
| for _ in range(self.power_iterations): |
| v.data = l2normalize(torch.mv(torch.t(w.view(height,-1).data), u.data)) |
| u.data = l2normalize(torch.mv(w.view(height,-1).data, v.data)) |
|
|
| sigma = u.dot(w.view(height, -1).mv(v)) |
| setattr(self.module, self.name, w / sigma.expand_as(w)) |
|
|
| def _made_params(self): |
| try: |
| u = getattr(self.module, self.name + "_u") |
| v = getattr(self.module, self.name + "_v") |
| w = getattr(self.module, self.name + "_bar") |
| return True |
| except AttributeError: |
| return False |
|
|
| def _make_params(self): |
| w = getattr(self.module, self.name) |
|
|
| height = w.data.shape[0] |
| width = w.view(height, -1).data.shape[1] |
|
|
| u = nn.Parameter(w.data.new(height).normal_(0, 1), requires_grad=False) |
| v = nn.Parameter(w.data.new(width).normal_(0, 1), requires_grad=False) |
| u.data = l2normalize(u.data) |
| v.data = l2normalize(v.data) |
| w_bar = nn.Parameter(w.data) |
|
|
| del self.module._parameters[self.name] |
|
|
| self.module.register_parameter(self.name + "_u", u) |
| self.module.register_parameter(self.name + "_v", v) |
| self.module.register_parameter(self.name + "_bar", w_bar) |
|
|
|
|
| def forward(self, *args): |
| self._update_u_v() |
| return self.module.forward(*args) |
|
|