mirror of
https://github.com/modelscope/modelscope.git
synced 2026-09-01 19:49:03 +02:00
新建模型 image_control_3d_portrait
Link: https://code.alibaba-inc.com/Ali-MaaS/MaaS-lib/codereview/14003131 * add image_control_3d_portrait code * update requirements * remove duplicate plyfile * update test_image_control_3d_portrait code * update test code * update requirements * change face landmark models and revise as cr suggests * add save images choice * add image_control_3d_portrait.jpg * mv ops related code to modelscope/ops
This commit is contained in:
@@ -125,6 +125,7 @@ class Models(object):
|
||||
image_try_on = 'image-try-on'
|
||||
human_image_generation = 'human-image-generation'
|
||||
image_view_transform = 'image-view-transform'
|
||||
image_control_3d_portrait = 'image-control-3d-portrait'
|
||||
|
||||
# nlp models
|
||||
bert = 'bert'
|
||||
@@ -447,6 +448,7 @@ class Pipelines(object):
|
||||
image_try_on = 'image-try-on'
|
||||
human_image_generation = 'human-image-generation'
|
||||
image_view_transform = 'image-view-transform'
|
||||
image_control_3d_portrait = 'image-control-3d-portrait'
|
||||
|
||||
# nlp tasks
|
||||
automatic_post_editing = 'automatic-post-editing'
|
||||
@@ -917,7 +919,10 @@ DEFAULT_MODEL_FOR_PIPELINE = {
|
||||
Tasks.human_image_generation: (Pipelines.human_image_generation,
|
||||
'damo/cv_FreqHPT_human-image-generation'),
|
||||
Tasks.image_view_transform: (Pipelines.image_view_transform,
|
||||
'damo/cv_image-view-transform')
|
||||
'damo/cv_image-view-transform'),
|
||||
Tasks.image_control_3d_portrait: (
|
||||
Pipelines.image_control_3d_portrait,
|
||||
'damo/cv_vit_image-control-3d-portrait-synthesis')
|
||||
}
|
||||
|
||||
|
||||
|
||||
22
modelscope/models/cv/image_control_3d_portrait/__init__.py
Normal file
22
modelscope/models/cv/image_control_3d_portrait/__init__.py
Normal file
@@ -0,0 +1,22 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from modelscope.utils.import_utils import LazyImportModule
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .image_control_3d_portrait import ImageControl3dPortrait
|
||||
|
||||
else:
|
||||
_import_structure = {
|
||||
'image_control_3d_portrait': ['ImageControl3dPortrait']
|
||||
}
|
||||
|
||||
import sys
|
||||
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
@@ -0,0 +1,468 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
from typing import Any, Dict
|
||||
|
||||
import cv2
|
||||
import json
|
||||
import numpy as np
|
||||
import PIL.Image as Image
|
||||
import torch
|
||||
import torchvision.transforms as transforms
|
||||
from scipy.io import loadmat
|
||||
|
||||
from modelscope.metainfo import Models
|
||||
from modelscope.models.base import Tensor, TorchModel
|
||||
from modelscope.models.builder import MODELS
|
||||
from modelscope.models.cv.face_detection.peppa_pig_face.facer import FaceAna
|
||||
from modelscope.utils.constant import ModelFile, Tasks
|
||||
from modelscope.utils.device import create_device
|
||||
from modelscope.utils.logger import get_logger
|
||||
from .network.camera_utils import FOV_to_intrinsics, LookAtPoseSampler
|
||||
from .network.shape_utils import convert_sdf_samples_to_ply
|
||||
from .network.triplane import TriPlaneGenerator
|
||||
from .network.triplane_encoder import TriplaneEncoder
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
__all__ = ['ImageControl3dPortrait']
|
||||
|
||||
|
||||
@MODELS.register_module(
|
||||
Tasks.image_control_3d_portrait,
|
||||
module_name=Models.image_control_3d_portrait)
|
||||
class ImageControl3dPortrait(TorchModel):
|
||||
|
||||
def __init__(self, model_dir: str, *args, **kwargs):
|
||||
"""initialize the image face fusion model from the `model_dir` path.
|
||||
|
||||
Args:
|
||||
model_dir (str): the model path.
|
||||
"""
|
||||
super().__init__(model_dir, *args, **kwargs)
|
||||
|
||||
logger.info('model params:{}'.format(kwargs))
|
||||
self.neural_rendering_resolution = kwargs[
|
||||
'neural_rendering_resolution']
|
||||
self.cam_radius = kwargs['cam_radius']
|
||||
self.fov_deg = kwargs['fov_deg']
|
||||
self.truncation_psi = kwargs['truncation_psi']
|
||||
self.truncation_cutoff = kwargs['truncation_cutoff']
|
||||
self.z_dim = kwargs['z_dim']
|
||||
self.image_size = kwargs['image_size']
|
||||
self.shape_res = kwargs['shape_res']
|
||||
self.pitch_range = kwargs['pitch_range']
|
||||
self.yaw_range = kwargs['yaw_range']
|
||||
self.max_batch = kwargs['max_batch']
|
||||
self.num_frames = kwargs['num_frames']
|
||||
self.box_warp = kwargs['box_warp']
|
||||
self.save_shape = kwargs['save_shape']
|
||||
self.save_images = kwargs['save_images']
|
||||
|
||||
device = kwargs['device']
|
||||
self.device = create_device(device)
|
||||
|
||||
self.facer = FaceAna(model_dir)
|
||||
|
||||
similarity_mat_path = os.path.join(model_dir, 'BFM',
|
||||
'similarity_Lm3D_all.mat')
|
||||
self.lm3d_std = self.load_lm3d(similarity_mat_path)
|
||||
|
||||
init_model_json = os.path.join(model_dir, 'configs',
|
||||
'init_encoder.json')
|
||||
with open(init_model_json, 'r') as fr:
|
||||
init_kwargs_encoder = json.load(fr)
|
||||
encoder_path = os.path.join(model_dir, ModelFile.TORCH_MODEL_FILE)
|
||||
self.model = TriplaneEncoder(**init_kwargs_encoder)
|
||||
ckpt_encoder = torch.load(encoder_path, map_location='cpu')
|
||||
model_state = self.convert_state_dict(ckpt_encoder['state_dict'])
|
||||
self.model.load_state_dict(model_state)
|
||||
self.model = self.model.to(self.device)
|
||||
self.model.eval()
|
||||
|
||||
init_args_G = ()
|
||||
init_netG_json = os.path.join(model_dir, 'configs', 'init_G.json')
|
||||
with open(init_netG_json, 'r') as fr:
|
||||
init_kwargs_G = json.load(fr)
|
||||
self.netG = TriPlaneGenerator(*init_args_G, **init_kwargs_G)
|
||||
netG_path = os.path.join(model_dir, 'ffhqrebalanced512-128.pth')
|
||||
ckpt_G = torch.load(netG_path)
|
||||
self.netG.load_state_dict(ckpt_G['G_ema'], strict=False)
|
||||
self.netG.neural_rendering_resolution = self.neural_rendering_resolution
|
||||
self.netG = self.netG.to(self.device)
|
||||
self.netG.eval()
|
||||
|
||||
self.intrinsics = FOV_to_intrinsics(self.fov_deg, device=self.device)
|
||||
col, row = np.meshgrid(
|
||||
np.arange(self.image_size), np.arange(self.image_size))
|
||||
np_coord = np.stack((col, row), axis=2) / self.image_size # [0,1]
|
||||
self.coord = torch.from_numpy(np_coord.astype(
|
||||
np.float32)).unsqueeze(0).permute(0, 3, 1, 2).to(self.device)
|
||||
|
||||
self.image_transform = transforms.Compose([
|
||||
transforms.Resize((self.image_size, self.image_size)),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])
|
||||
])
|
||||
|
||||
logger.info('init done')
|
||||
|
||||
def convert_state_dict(self, state_dict):
|
||||
if not next(iter(state_dict)).startswith('module.'):
|
||||
return state_dict
|
||||
new_state_dict = OrderedDict()
|
||||
|
||||
split_index = 0
|
||||
for cur_key, cur_value in state_dict.items():
|
||||
if cur_key.startswith('module.model'):
|
||||
split_index = 13
|
||||
elif cur_key.startswith('module'):
|
||||
split_index = 7
|
||||
|
||||
break
|
||||
|
||||
for k, v in state_dict.items():
|
||||
name = k[split_index:]
|
||||
new_state_dict[name] = v
|
||||
return new_state_dict
|
||||
|
||||
def detect_face(self, img):
|
||||
src_h, src_w, _ = img.shape
|
||||
boxes, landmarks, _ = self.facer.run(img)
|
||||
if boxes.shape[0] == 0:
|
||||
return None
|
||||
elif boxes.shape[0] > 1:
|
||||
max_area = 0
|
||||
max_index = 0
|
||||
for i in range(boxes.shape[0]):
|
||||
bbox_width = boxes[i][2] - boxes[i][0]
|
||||
bbox_height = boxes[i][3] - boxes[i][1]
|
||||
area = int(bbox_width) * int(bbox_height)
|
||||
if area > max_area:
|
||||
max_index = i
|
||||
max_area = area
|
||||
|
||||
return landmarks[max_index]
|
||||
else:
|
||||
return landmarks[0]
|
||||
|
||||
def get_f5p(self, landmarks, np_img):
|
||||
eye_left = self.find_pupil(landmarks[36:41], np_img)
|
||||
eye_right = self.find_pupil(landmarks[42:47], np_img)
|
||||
if eye_left is None or eye_right is None:
|
||||
logger.warning(
|
||||
'cannot find 5 points with find_pupil, used mean instead.!')
|
||||
eye_left = landmarks[36:41].mean(axis=0)
|
||||
eye_right = landmarks[42:47].mean(axis=0)
|
||||
nose = landmarks[30]
|
||||
mouth_left = landmarks[48]
|
||||
mouth_right = landmarks[54]
|
||||
f5p = [[eye_left[0], eye_left[1]], [eye_right[0], eye_right[1]],
|
||||
[nose[0], nose[1]], [mouth_left[0], mouth_left[1]],
|
||||
[mouth_right[0], mouth_right[1]]]
|
||||
return np.array(f5p)
|
||||
|
||||
def find_pupil(self, landmarks, np_img):
|
||||
h, w, _ = np_img.shape
|
||||
xmax = int(landmarks[:, 0].max())
|
||||
xmin = int(landmarks[:, 0].min())
|
||||
ymax = int(landmarks[:, 1].max())
|
||||
ymin = int(landmarks[:, 1].min())
|
||||
|
||||
if ymin >= ymax or xmin >= xmax or ymin < 0 or xmin < 0 or ymax > h or xmax > w:
|
||||
return None
|
||||
eye_img_bgr = np_img[ymin:ymax, xmin:xmax, :]
|
||||
eye_img = cv2.cvtColor(eye_img_bgr, cv2.COLOR_BGR2GRAY)
|
||||
eye_img = cv2.equalizeHist(eye_img)
|
||||
n_marks = landmarks - np.array([xmin, ymin]).reshape([1, 2])
|
||||
eye_mask = cv2.fillConvexPoly(
|
||||
np.zeros_like(eye_img), n_marks.astype(np.int32), 1)
|
||||
ret, thresh = cv2.threshold(eye_img, 100, 255,
|
||||
cv2.THRESH_BINARY | cv2.THRESH_OTSU)
|
||||
thresh = (1 - thresh / 255.) * eye_mask
|
||||
cnt = 0
|
||||
xm = []
|
||||
ym = []
|
||||
for i in range(thresh.shape[0]):
|
||||
for j in range(thresh.shape[1]):
|
||||
if thresh[i, j] > 0.5:
|
||||
xm.append(j)
|
||||
ym.append(i)
|
||||
cnt += 1
|
||||
if cnt != 0:
|
||||
xm.sort()
|
||||
ym.sort()
|
||||
xm = xm[cnt // 2]
|
||||
ym = ym[cnt // 2]
|
||||
else:
|
||||
xm = thresh.shape[1] / 2
|
||||
ym = thresh.shape[0] / 2
|
||||
|
||||
return xm + xmin, ym + ymin
|
||||
|
||||
def load_lm3d(self, similarity_mat_path):
|
||||
|
||||
Lm3D = loadmat(similarity_mat_path)
|
||||
Lm3D = Lm3D['lm']
|
||||
|
||||
lm_idx = np.array([31, 37, 40, 43, 46, 49, 55]) - 1
|
||||
lm_data1 = Lm3D[lm_idx[0], :]
|
||||
lm_data2 = np.mean(Lm3D[lm_idx[[1, 2]], :], 0)
|
||||
lm_data3 = np.mean(Lm3D[lm_idx[[3, 4]], :], 0)
|
||||
lm_data4 = Lm3D[lm_idx[5], :]
|
||||
lm_data5 = Lm3D[lm_idx[6], :]
|
||||
|
||||
Lm3D = np.stack([lm_data1, lm_data2, lm_data3, lm_data4, lm_data5],
|
||||
axis=0)
|
||||
|
||||
Lm3D = Lm3D[[1, 2, 0, 3, 4], :]
|
||||
|
||||
return Lm3D
|
||||
|
||||
def POS(self, xp, x):
|
||||
npts = xp.shape[1]
|
||||
|
||||
A = np.zeros([2 * npts, 8])
|
||||
|
||||
A[0:2 * npts - 1:2, 0:3] = x.transpose()
|
||||
A[0:2 * npts - 1:2, 3] = 1
|
||||
|
||||
A[1:2 * npts:2, 4:7] = x.transpose()
|
||||
A[1:2 * npts:2, 7] = 1
|
||||
|
||||
b = np.reshape(xp.transpose(), [2 * npts, 1])
|
||||
|
||||
k, _, _, _ = np.linalg.lstsq(A, b)
|
||||
|
||||
R1 = k[0:3]
|
||||
R2 = k[4:7]
|
||||
sTx = k[3]
|
||||
sTy = k[7]
|
||||
s = (np.linalg.norm(R1) + np.linalg.norm(R2)) / 2
|
||||
t = np.stack([sTx, sTy], axis=0)
|
||||
|
||||
return t, s
|
||||
|
||||
def resize_n_crop_img(self, img, lm, t, s, target_size=224., mask=None):
|
||||
w0, h0 = img.size
|
||||
w = (w0 * s).astype(np.int32)
|
||||
h = (h0 * s).astype(np.int32)
|
||||
left = (w / 2 - target_size / 2 + float(
|
||||
(t[0] - w0 / 2) * s)).astype(np.int32)
|
||||
right = left + target_size
|
||||
up = (h / 2 - target_size / 2 + float(
|
||||
(h0 / 2 - t[1]) * s)).astype(np.int32)
|
||||
below = up + target_size
|
||||
|
||||
img = img.resize((w, h), resample=Image.BICUBIC)
|
||||
img = img.crop((left, up, right, below))
|
||||
|
||||
if mask is not None:
|
||||
mask = mask.resize((w, h), resample=Image.BICUBIC)
|
||||
mask = mask.crop((left, up, right, below))
|
||||
|
||||
lm = np.stack([lm[:, 0] - t[0] + w0 / 2, lm[:, 1] - t[1] + h0 / 2],
|
||||
axis=1) * s
|
||||
lm = lm - np.reshape(
|
||||
np.array([(w / 2 - target_size / 2),
|
||||
(h / 2 - target_size / 2)]), [1, 2])
|
||||
|
||||
return img, lm, mask
|
||||
|
||||
def align_img(self,
|
||||
img,
|
||||
lm,
|
||||
lm3D,
|
||||
mask=None,
|
||||
target_size=224.,
|
||||
rescale_factor=102.):
|
||||
w0, h0 = img.size
|
||||
lm5p = lm
|
||||
t, s = self.POS(lm5p.transpose(), lm3D.transpose())
|
||||
s = rescale_factor / s
|
||||
|
||||
img_new, lm_new, mask_new = self.resize_n_crop_img(
|
||||
img, lm, t, s, target_size=target_size, mask=mask)
|
||||
trans_params = np.array([w0, h0, s, t[0], t[1]], dtype=object)
|
||||
|
||||
return trans_params, img_new, lm_new, mask_new
|
||||
|
||||
def crop_image(self, img, lm):
|
||||
_, H = img.size
|
||||
lm[:, -1] = H - 1 - lm[:, -1]
|
||||
|
||||
target_size = 1024.
|
||||
rescale_factor = 300
|
||||
center_crop_size = 700
|
||||
output_size = 512
|
||||
|
||||
_, im_high, _, _, = self.align_img(
|
||||
img,
|
||||
lm,
|
||||
self.lm3d_std,
|
||||
target_size=target_size,
|
||||
rescale_factor=rescale_factor)
|
||||
|
||||
left = int(im_high.size[0] / 2 - center_crop_size / 2)
|
||||
upper = int(im_high.size[1] / 2 - center_crop_size / 2)
|
||||
right = left + center_crop_size
|
||||
lower = upper + center_crop_size
|
||||
im_cropped = im_high.crop((left, upper, right, lower))
|
||||
im_cropped = im_cropped.resize((output_size, output_size),
|
||||
resample=Image.LANCZOS)
|
||||
logger.info('crop image done!')
|
||||
return im_cropped
|
||||
|
||||
def create_samples(self, N=256, voxel_origin=[0, 0, 0], cube_length=2.0):
|
||||
voxel_origin = np.array(voxel_origin) - cube_length / 2
|
||||
voxel_size = cube_length / (N - 1)
|
||||
|
||||
overall_index = torch.arange(0, N**3, 1, out=torch.LongTensor())
|
||||
samples = torch.zeros(N**3, 3)
|
||||
|
||||
samples[:, 2] = overall_index % N
|
||||
samples[:, 1] = (overall_index.float() / N) % N
|
||||
samples[:, 0] = ((overall_index.float() / N) / N) % N
|
||||
|
||||
samples[:, 0] = (samples[:, 0] * voxel_size) + voxel_origin[2]
|
||||
samples[:, 1] = (samples[:, 1] * voxel_size) + voxel_origin[1]
|
||||
samples[:, 2] = (samples[:, 2] * voxel_size) + voxel_origin[0]
|
||||
|
||||
return samples.unsqueeze(0), voxel_origin, voxel_size
|
||||
|
||||
def numpy_array_to_video(self, numpy_list, video_out_path):
|
||||
assert len(numpy_list) > 0
|
||||
video_height = numpy_list[0].shape[0]
|
||||
video_width = numpy_list[0].shape[1]
|
||||
|
||||
out_video_size = (video_width, video_height)
|
||||
output_video_fourcc = cv2.VideoWriter_fourcc('m', 'p', '4', 'v')
|
||||
video_write_capture = cv2.VideoWriter(video_out_path,
|
||||
output_video_fourcc, 30,
|
||||
out_video_size)
|
||||
|
||||
for frame in numpy_list:
|
||||
frame_bgr = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
video_write_capture.write(frame_bgr)
|
||||
|
||||
video_write_capture.release()
|
||||
|
||||
def inference(self, image_path, save_dir):
|
||||
basename = os.path.basename(image_path).split('.')[0]
|
||||
img = Image.open(image_path).convert('RGB')
|
||||
img_array = np.array(img)
|
||||
img_bgr = img_array[:, :, ::-1]
|
||||
landmark = self.detect_face(img_array)
|
||||
if landmark is None:
|
||||
logger.warning('No face detected in the image!')
|
||||
f5p = self.get_f5p(landmark, img_bgr)
|
||||
|
||||
logger.info('f5p is:{}'.format(f5p))
|
||||
img_cropped = self.crop_image(img, f5p)
|
||||
img_cropped.save(os.path.join(save_dir, 'crop.jpg'))
|
||||
|
||||
in_image = self.image_transform(img_cropped).unsqueeze(0).to(
|
||||
self.device)
|
||||
input = torch.cat((in_image, self.coord), 1)
|
||||
|
||||
save_video_path = os.path.join(save_dir, f'{basename}.mp4')
|
||||
pred_imgs = []
|
||||
|
||||
for frame_idx in range(self.num_frames):
|
||||
cam_pivot = torch.tensor([0, 0, 0.2], device=self.device)
|
||||
|
||||
cam2world_pose = LookAtPoseSampler.sample(
|
||||
3.14 / 2 + self.yaw_range
|
||||
* np.sin(2 * 3.14 * frame_idx / self.num_frames),
|
||||
3.14 / 2 - 0.05 + self.pitch_range
|
||||
* np.cos(2 * 3.14 * frame_idx / self.num_frames),
|
||||
cam_pivot,
|
||||
radius=self.cam_radius,
|
||||
device=self.device)
|
||||
|
||||
camera_params = torch.cat([
|
||||
cam2world_pose.reshape(-1, 16),
|
||||
self.intrinsics.reshape(-1, 9)
|
||||
], 1)
|
||||
|
||||
conditioning_cam2world_pose = LookAtPoseSampler.sample(
|
||||
np.pi / 2,
|
||||
np.pi / 2,
|
||||
cam_pivot,
|
||||
radius=self.cam_radius,
|
||||
device=self.device)
|
||||
conditioning_params = torch.cat([
|
||||
conditioning_cam2world_pose.reshape(-1, 16),
|
||||
self.intrinsics.reshape(-1, 9)
|
||||
], 1)
|
||||
|
||||
z = torch.from_numpy(np.random.randn(1,
|
||||
self.z_dim)).to(self.device)
|
||||
|
||||
with torch.no_grad():
|
||||
ws = self.netG.mapping(
|
||||
z,
|
||||
conditioning_params,
|
||||
truncation_psi=self.truncation_psi,
|
||||
truncation_cutoff=self.truncation_cutoff)
|
||||
|
||||
planes, pred_depth, pred_feature, pred_rgb, pred_sr, _, _, _, _ = self.model(
|
||||
ws, input, camera_params, None)
|
||||
|
||||
pred_img = (pred_sr.permute(0, 2, 3, 1) * 127.5 + 128).clamp(
|
||||
0, 255).to(torch.uint8)
|
||||
pred_img = pred_img.squeeze().cpu().numpy()
|
||||
if self.save_images:
|
||||
cv2.imwrite(
|
||||
os.path.join(save_dir, '{}.jpg'.format(frame_idx)),
|
||||
pred_img[:, :, ::-1])
|
||||
pred_imgs.append(pred_img)
|
||||
|
||||
self.numpy_array_to_video(pred_imgs, save_video_path)
|
||||
|
||||
if self.save_shape:
|
||||
max_batch = 1000000
|
||||
|
||||
samples, voxel_origin, voxel_size = self.create_samples(
|
||||
N=self.shape_res,
|
||||
voxel_origin=[0, 0, 0],
|
||||
cube_length=self.box_warp)
|
||||
samples = samples.to(z.device)
|
||||
sigmas = torch.zeros((samples.shape[0], samples.shape[1], 1),
|
||||
device=z.device)
|
||||
transformed_ray_directions_expanded = torch.zeros(
|
||||
(samples.shape[0], max_batch, 3), device=z.device)
|
||||
transformed_ray_directions_expanded[..., -1] = -1
|
||||
|
||||
head = 0
|
||||
with torch.no_grad():
|
||||
while head < samples.shape[1]:
|
||||
torch.manual_seed(0)
|
||||
sigma = self.model.sample(
|
||||
samples[:, head:head + max_batch],
|
||||
transformed_ray_directions_expanded[:, :samples.
|
||||
shape[1] - head],
|
||||
planes)['sigma']
|
||||
sigmas[:, head:head + max_batch] = sigma
|
||||
head += max_batch
|
||||
|
||||
sigmas = sigmas.reshape((self.shape_res, self.shape_res,
|
||||
self.shape_res)).cpu().numpy()
|
||||
sigmas = np.flip(sigmas, 0)
|
||||
|
||||
pad = int(30 * self.shape_res / 256)
|
||||
pad_value = -1000
|
||||
sigmas[:pad] = pad_value
|
||||
sigmas[-pad:] = pad_value
|
||||
sigmas[:, :pad] = pad_value
|
||||
sigmas[:, -pad:] = pad_value
|
||||
sigmas[:, :, :pad] = pad_value
|
||||
sigmas[:, :, -pad:] = pad_value
|
||||
convert_sdf_samples_to_ply(
|
||||
np.transpose(sigmas, (2, 1, 0)), [0, 0, 0],
|
||||
1,
|
||||
os.path.join(save_dir, f'{basename}.ply'),
|
||||
level=10)
|
||||
|
||||
logger.info('model inference done')
|
||||
@@ -0,0 +1,195 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""
|
||||
Helper functions for constructing camera parameter matrices. Primarily used in visualization and inference scripts.
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from .volumetric_rendering import math_utils
|
||||
|
||||
|
||||
class GaussianCameraPoseSampler:
|
||||
"""
|
||||
Samples pitch and yaw from a Gaussian distribution and returns a camera pose.
|
||||
Camera is specified as looking at the origin.
|
||||
If horizontal and vertical stddev (specified in radians) are zero, gives a
|
||||
deterministic camera pose with yaw=horizontal_mean, pitch=vertical_mean.
|
||||
The coordinate system is specified with y-up, z-forward, x-left.
|
||||
Horizontal mean is the azimuthal angle (rotation around y axis) in radians,
|
||||
vertical mean is the polar angle (angle from the y axis) in radians.
|
||||
A point along the z-axis has azimuthal_angle=0, polar_angle=pi/2.
|
||||
|
||||
Example:
|
||||
For a camera pose looking at the origin with the camera at position [0, 0, 1]:
|
||||
cam2world = GaussianCameraPoseSampler.sample(math.pi/2, math.pi/2, radius=1)
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def sample(horizontal_mean,
|
||||
vertical_mean,
|
||||
horizontal_stddev=0,
|
||||
vertical_stddev=0,
|
||||
radius=1,
|
||||
batch_size=1,
|
||||
device='cpu'):
|
||||
h = torch.randn((batch_size, 1),
|
||||
device=device) * horizontal_stddev + horizontal_mean
|
||||
v = torch.randn(
|
||||
(batch_size, 1), device=device) * vertical_stddev + vertical_mean
|
||||
v = torch.clamp(v, 1e-5, math.pi - 1e-5)
|
||||
|
||||
theta = h
|
||||
v = v / math.pi
|
||||
phi = torch.arccos(1 - 2 * v)
|
||||
|
||||
camera_origins = torch.zeros((batch_size, 3), device=device)
|
||||
|
||||
camera_origins[:, 0:1] = radius * torch.sin(phi) * torch.cos(math.pi
|
||||
- theta)
|
||||
camera_origins[:, 2:3] = radius * torch.sin(phi) * torch.sin(math.pi
|
||||
- theta)
|
||||
camera_origins[:, 1:2] = radius * torch.cos(phi)
|
||||
|
||||
forward_vectors = math_utils.normalize_vecs(-camera_origins)
|
||||
return create_cam2world_matrix(forward_vectors, camera_origins)
|
||||
|
||||
|
||||
class LookAtPoseSampler:
|
||||
"""
|
||||
Same as GaussianCameraPoseSampler, except the
|
||||
camera is specified as looking at 'lookat_position', a 3-vector.
|
||||
|
||||
Example:
|
||||
For a camera pose looking at the origin with the camera at position [0, 0, 1]:
|
||||
cam2world = LookAtPoseSampler.sample(math.pi/2, math.pi/2, torch.tensor([0, 0, 0]), radius=1)
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def sample(horizontal_mean,
|
||||
vertical_mean,
|
||||
lookat_position,
|
||||
horizontal_stddev=0,
|
||||
vertical_stddev=0,
|
||||
radius=1,
|
||||
batch_size=1,
|
||||
device='cpu'):
|
||||
h = torch.randn((batch_size, 1),
|
||||
device=device) * horizontal_stddev + horizontal_mean
|
||||
v = torch.randn(
|
||||
(batch_size, 1), device=device) * vertical_stddev + vertical_mean
|
||||
v = torch.clamp(v, 1e-5, math.pi - 1e-5)
|
||||
|
||||
theta = h
|
||||
v = v / math.pi
|
||||
phi = torch.arccos(1 - 2 * v)
|
||||
|
||||
camera_origins = torch.zeros((batch_size, 3), device=device)
|
||||
|
||||
camera_origins[:, 0:1] = radius * torch.sin(phi) * torch.cos(math.pi
|
||||
- theta)
|
||||
camera_origins[:, 2:3] = radius * torch.sin(phi) * torch.sin(math.pi
|
||||
- theta)
|
||||
camera_origins[:, 1:2] = radius * torch.cos(phi)
|
||||
|
||||
# forward_vectors = math_utils.normalize_vecs(-camera_origins)
|
||||
forward_vectors = math_utils.normalize_vecs(lookat_position
|
||||
- camera_origins)
|
||||
return create_cam2world_matrix(forward_vectors, camera_origins)
|
||||
|
||||
|
||||
class UniformCameraPoseSampler:
|
||||
"""
|
||||
Same as GaussianCameraPoseSampler, except the
|
||||
pose is sampled from a uniform distribution with range +-[horizontal/vertical]_stddev.
|
||||
|
||||
Example:
|
||||
For a batch of random camera poses looking at the origin with yaw sampled from [-pi/2, +pi/2] radians:
|
||||
|
||||
cam2worlds = UniformCameraPoseSampler.sample
|
||||
(math.pi/2, math.pi/2, horizontal_stddev=math.pi/2, radius=1, batch_size=16)
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def sample(horizontal_mean,
|
||||
vertical_mean,
|
||||
horizontal_stddev=0,
|
||||
vertical_stddev=0,
|
||||
radius=1,
|
||||
batch_size=1,
|
||||
device='cpu'):
|
||||
h = (torch.rand((batch_size, 1), device=device) * 2
|
||||
- 1) * horizontal_stddev + horizontal_mean
|
||||
v = (torch.rand((batch_size, 1), device=device) * 2
|
||||
- 1) * vertical_stddev + vertical_mean
|
||||
v = torch.clamp(v, 1e-5, math.pi - 1e-5)
|
||||
|
||||
theta = h
|
||||
v = v / math.pi
|
||||
phi = torch.arccos(1 - 2 * v)
|
||||
|
||||
camera_origins = torch.zeros((batch_size, 3), device=device)
|
||||
|
||||
camera_origins[:, 0:1] = radius * torch.sin(phi) * torch.cos(math.pi
|
||||
- theta)
|
||||
camera_origins[:, 2:3] = radius * torch.sin(phi) * torch.sin(math.pi
|
||||
- theta)
|
||||
camera_origins[:, 1:2] = radius * torch.cos(phi)
|
||||
|
||||
forward_vectors = math_utils.normalize_vecs(-camera_origins)
|
||||
return create_cam2world_matrix(forward_vectors, camera_origins)
|
||||
|
||||
|
||||
def create_cam2world_matrix(forward_vector, origin):
|
||||
"""
|
||||
Takes in the direction the camera is pointing and the camera origin and returns a cam2world matrix.
|
||||
Works on batches of forward_vectors, origins. Assumes y-axis is up and that there is no camera roll.
|
||||
"""
|
||||
|
||||
forward_vector = math_utils.normalize_vecs(forward_vector)
|
||||
up_vector = torch.tensor([0, 1, 0],
|
||||
dtype=torch.float,
|
||||
device=origin.device).expand_as(forward_vector)
|
||||
|
||||
right_vector = -math_utils.normalize_vecs(
|
||||
torch.cross(up_vector, forward_vector, dim=-1))
|
||||
up_vector = math_utils.normalize_vecs(
|
||||
torch.cross(forward_vector, right_vector, dim=-1))
|
||||
|
||||
rotation_matrix = torch.eye(
|
||||
4, device=origin.device).unsqueeze(0).repeat(forward_vector.shape[0],
|
||||
1, 1)
|
||||
rotation_matrix[:, :3, :3] = torch.stack(
|
||||
(right_vector, up_vector, forward_vector), axis=-1)
|
||||
|
||||
translation_matrix = torch.eye(
|
||||
4, device=origin.device).unsqueeze(0).repeat(forward_vector.shape[0],
|
||||
1, 1)
|
||||
translation_matrix[:, :3, 3] = origin
|
||||
cam2world = (translation_matrix @ rotation_matrix)[:, :, :]
|
||||
assert (cam2world.shape[1:] == (4, 4))
|
||||
return cam2world
|
||||
|
||||
|
||||
def FOV_to_intrinsics(fov_degrees, device='cpu'):
|
||||
"""
|
||||
Creates a 3x3 camera intrinsics matrix from the camera field of view, specified in degrees.
|
||||
Note the intrinsics are returned as normalized by image size, rather than in pixel units.
|
||||
Assumes principal point is at image center.
|
||||
"""
|
||||
|
||||
focal_length = float(1 / (math.tan(fov_degrees * 3.14159 / 360) * 1.414))
|
||||
intrinsics = torch.tensor(
|
||||
[[focal_length, 0, 0.5], [0, focal_length, 0.5], [0, 0, 1]],
|
||||
device=device)
|
||||
return intrinsics
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,65 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""
|
||||
Utils for extracting 3D shapes using marching cubes. Based on code from DeepSDF (Park et al.)
|
||||
|
||||
Takes as input an .mrc file and extracts a mesh.
|
||||
|
||||
Ex.
|
||||
python shape_utils.py my_shape.mrc
|
||||
Ex.
|
||||
python shape_utils.py myshapes_directory --level=12
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
import plyfile
|
||||
import skimage.measure
|
||||
|
||||
|
||||
def convert_sdf_samples_to_ply(numpy_3d_sdf_tensor,
|
||||
voxel_grid_origin,
|
||||
voxel_size,
|
||||
ply_filename_out,
|
||||
offset=None,
|
||||
scale=None,
|
||||
level=0.0):
|
||||
|
||||
verts, faces, normals, values = skimage.measure.marching_cubes(
|
||||
numpy_3d_sdf_tensor, level=level, spacing=[voxel_size] * 3)
|
||||
mesh_points = np.zeros_like(verts)
|
||||
mesh_points[:, 0] = voxel_grid_origin[0] + verts[:, 0]
|
||||
mesh_points[:, 1] = voxel_grid_origin[1] + verts[:, 1]
|
||||
mesh_points[:, 2] = voxel_grid_origin[2] + verts[:, 2]
|
||||
|
||||
if scale is not None:
|
||||
mesh_points = mesh_points / scale
|
||||
if offset is not None:
|
||||
mesh_points = mesh_points - offset
|
||||
|
||||
num_verts = verts.shape[0]
|
||||
num_faces = faces.shape[0]
|
||||
|
||||
verts_tuple = np.zeros((num_verts, ),
|
||||
dtype=[('x', 'f4'), ('y', 'f4'), ('z', 'f4')])
|
||||
|
||||
for i in range(0, num_verts):
|
||||
verts_tuple[i] = tuple(mesh_points[i, :])
|
||||
|
||||
faces_building = []
|
||||
for i in range(0, num_faces):
|
||||
faces_building.append(((faces[i, :].tolist(), )))
|
||||
faces_tuple = np.array(
|
||||
faces_building, dtype=[('vertex_indices', 'i4', (3, ))])
|
||||
|
||||
el_verts = plyfile.PlyElement.describe(verts_tuple, 'vertex')
|
||||
el_faces = plyfile.PlyElement.describe(faces_tuple, 'face')
|
||||
|
||||
ply_data = plyfile.PlyData([el_verts, el_faces])
|
||||
ply_data.write(ply_filename_out)
|
||||
@@ -0,0 +1,493 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""Superresolution network architectures from the paper
|
||||
"Efficient Geometry-aware 3D Generative Adversarial Networks"."""
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from modelscope.ops.image_control_3d_portrait.torch_utils import (misc,
|
||||
persistence)
|
||||
from modelscope.ops.image_control_3d_portrait.torch_utils.ops import upfirdn2d
|
||||
from .networks_stylegan2 import (Conv2dLayer, SynthesisBlock, SynthesisLayer,
|
||||
ToRGBLayer)
|
||||
|
||||
|
||||
# for 512x512 generation
|
||||
@persistence.persistent_class
|
||||
class SuperresolutionHybrid8X(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
channels,
|
||||
img_resolution,
|
||||
sr_num_fp16_res,
|
||||
sr_antialias,
|
||||
num_fp16_res=4,
|
||||
conv_clamp=None,
|
||||
channel_base=None,
|
||||
channel_max=None, # IGNORE
|
||||
**block_kwargs):
|
||||
super().__init__()
|
||||
assert img_resolution == 512
|
||||
|
||||
use_fp16 = sr_num_fp16_res > 0
|
||||
self.input_resolution = 128
|
||||
self.sr_antialias = sr_antialias
|
||||
self.block0 = SynthesisBlock(
|
||||
channels,
|
||||
128,
|
||||
w_dim=512,
|
||||
resolution=256,
|
||||
img_channels=3,
|
||||
is_last=False,
|
||||
use_fp16=use_fp16,
|
||||
conv_clamp=(256 if use_fp16 else None),
|
||||
**block_kwargs)
|
||||
self.block1 = SynthesisBlock(
|
||||
128,
|
||||
64,
|
||||
w_dim=512,
|
||||
resolution=512,
|
||||
img_channels=3,
|
||||
is_last=True,
|
||||
use_fp16=use_fp16,
|
||||
conv_clamp=(256 if use_fp16 else None),
|
||||
**block_kwargs)
|
||||
self.register_buffer('resample_filter',
|
||||
upfirdn2d.setup_filter([1, 3, 3, 1]))
|
||||
|
||||
def forward(self, rgb, x, ws, **block_kwargs):
|
||||
ws = ws[:, -1:, :].repeat(1, 3, 1)
|
||||
|
||||
if x.shape[-1] != self.input_resolution:
|
||||
x = torch.nn.functional.interpolate(
|
||||
x,
|
||||
size=(self.input_resolution, self.input_resolution),
|
||||
mode='bilinear',
|
||||
align_corners=False,
|
||||
antialias=self.sr_antialias)
|
||||
rgb = torch.nn.functional.interpolate(
|
||||
rgb,
|
||||
size=(self.input_resolution, self.input_resolution),
|
||||
mode='bilinear',
|
||||
align_corners=False,
|
||||
antialias=self.sr_antialias)
|
||||
|
||||
x, rgb = self.block0(x, rgb, ws, **block_kwargs)
|
||||
x, rgb = self.block1(x, rgb, ws, **block_kwargs)
|
||||
return rgb
|
||||
|
||||
|
||||
# for 256x256 generation
|
||||
@persistence.persistent_class
|
||||
class SuperresolutionHybrid4X(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
channels,
|
||||
img_resolution,
|
||||
sr_num_fp16_res,
|
||||
sr_antialias,
|
||||
num_fp16_res=4,
|
||||
conv_clamp=None,
|
||||
channel_base=None,
|
||||
channel_max=None, # IGNORE
|
||||
**block_kwargs):
|
||||
super().__init__()
|
||||
assert img_resolution == 256
|
||||
use_fp16 = sr_num_fp16_res > 0
|
||||
self.sr_antialias = sr_antialias
|
||||
self.input_resolution = 128
|
||||
self.block0 = SynthesisBlockNoUp(
|
||||
channels,
|
||||
128,
|
||||
w_dim=512,
|
||||
resolution=128,
|
||||
img_channels=3,
|
||||
is_last=False,
|
||||
use_fp16=use_fp16,
|
||||
conv_clamp=(256 if use_fp16 else None),
|
||||
**block_kwargs)
|
||||
self.block1 = SynthesisBlock(
|
||||
128,
|
||||
64,
|
||||
w_dim=512,
|
||||
resolution=256,
|
||||
img_channels=3,
|
||||
is_last=True,
|
||||
use_fp16=use_fp16,
|
||||
conv_clamp=(256 if use_fp16 else None),
|
||||
**block_kwargs)
|
||||
self.register_buffer('resample_filter',
|
||||
upfirdn2d.setup_filter([1, 3, 3, 1]))
|
||||
|
||||
def forward(self, rgb, x, ws, **block_kwargs):
|
||||
ws = ws[:, -1:, :].repeat(1, 3, 1)
|
||||
|
||||
if x.shape[-1] < self.input_resolution:
|
||||
x = torch.nn.functional.interpolate(
|
||||
x,
|
||||
size=(self.input_resolution, self.input_resolution),
|
||||
mode='bilinear',
|
||||
align_corners=False,
|
||||
antialias=self.sr_antialias)
|
||||
rgb = torch.nn.functional.interpolate(
|
||||
rgb,
|
||||
size=(self.input_resolution, self.input_resolution),
|
||||
mode='bilinear',
|
||||
align_corners=False,
|
||||
antialias=self.sr_antialias)
|
||||
|
||||
x, rgb = self.block0(x, rgb, ws, **block_kwargs)
|
||||
x, rgb = self.block1(x, rgb, ws, **block_kwargs)
|
||||
return rgb
|
||||
|
||||
|
||||
# for 128 x 128 generation
|
||||
@persistence.persistent_class
|
||||
class SuperresolutionHybrid2X(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
channels,
|
||||
img_resolution,
|
||||
sr_num_fp16_res,
|
||||
sr_antialias,
|
||||
num_fp16_res=4,
|
||||
conv_clamp=None,
|
||||
channel_base=None,
|
||||
channel_max=None, # IGNORE
|
||||
**block_kwargs):
|
||||
super().__init__()
|
||||
assert img_resolution == 128
|
||||
|
||||
use_fp16 = sr_num_fp16_res > 0
|
||||
self.input_resolution = 64
|
||||
self.sr_antialias = sr_antialias
|
||||
self.block0 = SynthesisBlockNoUp(
|
||||
channels,
|
||||
128,
|
||||
w_dim=512,
|
||||
resolution=64,
|
||||
img_channels=3,
|
||||
is_last=False,
|
||||
use_fp16=use_fp16,
|
||||
conv_clamp=(256 if use_fp16 else None),
|
||||
**block_kwargs)
|
||||
self.block1 = SynthesisBlock(
|
||||
128,
|
||||
64,
|
||||
w_dim=512,
|
||||
resolution=128,
|
||||
img_channels=3,
|
||||
is_last=True,
|
||||
use_fp16=use_fp16,
|
||||
conv_clamp=(256 if use_fp16 else None),
|
||||
**block_kwargs)
|
||||
self.register_buffer('resample_filter',
|
||||
upfirdn2d.setup_filter([1, 3, 3, 1]))
|
||||
|
||||
def forward(self, rgb, x, ws, **block_kwargs):
|
||||
ws = ws[:, -1:, :].repeat(1, 3, 1)
|
||||
|
||||
if x.shape[-1] != self.input_resolution:
|
||||
x = torch.nn.functional.interpolate(
|
||||
x,
|
||||
size=(self.input_resolution, self.input_resolution),
|
||||
mode='bilinear',
|
||||
align_corners=False,
|
||||
antialias=self.sr_antialias)
|
||||
rgb = torch.nn.functional.interpolate(
|
||||
rgb,
|
||||
size=(self.input_resolution, self.input_resolution),
|
||||
mode='bilinear',
|
||||
align_corners=False,
|
||||
antialias=self.sr_antialias)
|
||||
|
||||
x, rgb = self.block0(x, rgb, ws, **block_kwargs)
|
||||
x, rgb = self.block1(x, rgb, ws, **block_kwargs)
|
||||
return rgb
|
||||
|
||||
|
||||
# TODO: Delete (here for backwards compatibility with old 256x256 models)
|
||||
@persistence.persistent_class
|
||||
class SuperresolutionHybridDeepfp32(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
channels,
|
||||
img_resolution,
|
||||
sr_num_fp16_res,
|
||||
num_fp16_res=4,
|
||||
conv_clamp=None,
|
||||
channel_base=None,
|
||||
channel_max=None, # IGNORE
|
||||
**block_kwargs):
|
||||
super().__init__()
|
||||
assert img_resolution == 256
|
||||
use_fp16 = sr_num_fp16_res > 0
|
||||
|
||||
self.input_resolution = 128
|
||||
self.block0 = SynthesisBlockNoUp(
|
||||
channels,
|
||||
128,
|
||||
w_dim=512,
|
||||
resolution=128,
|
||||
img_channels=3,
|
||||
is_last=False,
|
||||
use_fp16=use_fp16,
|
||||
conv_clamp=(256 if use_fp16 else None),
|
||||
**block_kwargs)
|
||||
self.block1 = SynthesisBlock(
|
||||
128,
|
||||
64,
|
||||
w_dim=512,
|
||||
resolution=256,
|
||||
img_channels=3,
|
||||
is_last=True,
|
||||
use_fp16=use_fp16,
|
||||
conv_clamp=(256 if use_fp16 else None),
|
||||
**block_kwargs)
|
||||
self.register_buffer('resample_filter',
|
||||
upfirdn2d.setup_filter([1, 3, 3, 1]))
|
||||
|
||||
def forward(self, rgb, x, ws, **block_kwargs):
|
||||
ws = ws[:, -1:, :].repeat(1, 3, 1)
|
||||
|
||||
if x.shape[-1] < self.input_resolution:
|
||||
x = torch.nn.functional.interpolate(
|
||||
x,
|
||||
size=(self.input_resolution, self.input_resolution),
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
rgb = torch.nn.functional.interpolate(
|
||||
rgb,
|
||||
size=(self.input_resolution, self.input_resolution),
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
|
||||
x, rgb = self.block0(x, rgb, ws, **block_kwargs)
|
||||
x, rgb = self.block1(x, rgb, ws, **block_kwargs)
|
||||
return rgb
|
||||
|
||||
|
||||
@persistence.persistent_class
|
||||
class SynthesisBlockNoUp(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels, # Number of input channels, 0 = first block.
|
||||
out_channels, # Number of output channels.
|
||||
w_dim, # Intermediate latent (W) dimensionality.
|
||||
resolution, # Resolution of this block.
|
||||
img_channels, # Number of output color channels.
|
||||
is_last, # Is this the last block?
|
||||
architecture='skip', # Architecture: 'orig', 'skip', 'resnet'.
|
||||
resample_filter=[
|
||||
1, 3, 3, 1
|
||||
], # Low-pass filter to apply when resampling activations.
|
||||
conv_clamp=256, # Clamp the output of convolution layers to +-X, None = disable clamping.
|
||||
use_fp16=False, # Use FP16 for this block?
|
||||
fp16_channels_last=False, # Use channels-last memory format with FP16?
|
||||
fused_modconv_default=True, # Default value of fused_modconv.
|
||||
**layer_kwargs, # Arguments for SynthesisLayer.
|
||||
):
|
||||
assert architecture in ['orig', 'skip', 'resnet']
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.w_dim = w_dim
|
||||
self.resolution = resolution
|
||||
self.img_channels = img_channels
|
||||
self.is_last = is_last
|
||||
self.architecture = architecture
|
||||
self.use_fp16 = use_fp16
|
||||
self.channels_last = (use_fp16 and fp16_channels_last)
|
||||
self.fused_modconv_default = fused_modconv_default
|
||||
self.register_buffer('resample_filter',
|
||||
upfirdn2d.setup_filter(resample_filter))
|
||||
self.num_conv = 0
|
||||
self.num_torgb = 0
|
||||
|
||||
if in_channels == 0:
|
||||
self.const = torch.nn.Parameter(
|
||||
torch.randn([out_channels, resolution, resolution]))
|
||||
|
||||
if in_channels != 0:
|
||||
self.conv0 = SynthesisLayer(
|
||||
in_channels,
|
||||
out_channels,
|
||||
w_dim=w_dim,
|
||||
resolution=resolution,
|
||||
conv_clamp=conv_clamp,
|
||||
channels_last=self.channels_last,
|
||||
**layer_kwargs)
|
||||
self.num_conv += 1
|
||||
|
||||
self.conv1 = SynthesisLayer(
|
||||
out_channels,
|
||||
out_channels,
|
||||
w_dim=w_dim,
|
||||
resolution=resolution,
|
||||
conv_clamp=conv_clamp,
|
||||
channels_last=self.channels_last,
|
||||
**layer_kwargs)
|
||||
self.num_conv += 1
|
||||
|
||||
if is_last or architecture == 'skip':
|
||||
self.torgb = ToRGBLayer(
|
||||
out_channels,
|
||||
img_channels,
|
||||
w_dim=w_dim,
|
||||
conv_clamp=conv_clamp,
|
||||
channels_last=self.channels_last)
|
||||
self.num_torgb += 1
|
||||
|
||||
if in_channels != 0 and architecture == 'resnet':
|
||||
self.skip = Conv2dLayer(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size=1,
|
||||
bias=False,
|
||||
up=2,
|
||||
resample_filter=resample_filter,
|
||||
channels_last=self.channels_last)
|
||||
|
||||
def forward(self,
|
||||
x,
|
||||
img,
|
||||
ws,
|
||||
force_fp32=False,
|
||||
fused_modconv=None,
|
||||
update_emas=False,
|
||||
**layer_kwargs):
|
||||
_ = update_emas # unused
|
||||
misc.assert_shape(ws,
|
||||
[None, self.num_conv + self.num_torgb, self.w_dim])
|
||||
w_iter = iter(ws.unbind(dim=1))
|
||||
if ws.device.type != 'cuda':
|
||||
force_fp32 = True
|
||||
dtype = torch.float16 if self.use_fp16 and not force_fp32 else torch.float32
|
||||
memory_format = torch.channels_last if self.channels_last and not force_fp32 else torch.contiguous_format
|
||||
if fused_modconv is None:
|
||||
fused_modconv = self.fused_modconv_default
|
||||
if fused_modconv == 'inference_only':
|
||||
fused_modconv = (not self.training)
|
||||
|
||||
# Input.
|
||||
if self.in_channels == 0:
|
||||
x = self.const.to(dtype=dtype, memory_format=memory_format)
|
||||
x = x.unsqueeze(0).repeat([ws.shape[0], 1, 1, 1])
|
||||
else:
|
||||
misc.assert_shape(
|
||||
x, [None, self.in_channels, self.resolution, self.resolution])
|
||||
x = x.to(dtype=dtype, memory_format=memory_format)
|
||||
|
||||
# Main layers.
|
||||
if self.in_channels == 0:
|
||||
x = self.conv1(
|
||||
x, next(w_iter), fused_modconv=fused_modconv, **layer_kwargs)
|
||||
elif self.architecture == 'resnet':
|
||||
y = self.skip(x, gain=np.sqrt(0.5))
|
||||
x = self.conv0(
|
||||
x, next(w_iter), fused_modconv=fused_modconv, **layer_kwargs)
|
||||
x = self.conv1(
|
||||
x,
|
||||
next(w_iter),
|
||||
fused_modconv=fused_modconv,
|
||||
gain=np.sqrt(0.5),
|
||||
**layer_kwargs)
|
||||
x = y.add_(x)
|
||||
else:
|
||||
x = self.conv0(
|
||||
x, next(w_iter), fused_modconv=fused_modconv, **layer_kwargs)
|
||||
x = self.conv1(
|
||||
x, next(w_iter), fused_modconv=fused_modconv, **layer_kwargs)
|
||||
|
||||
# ToRGB.
|
||||
# if img is not None:
|
||||
# misc.assert_shape(img, [None, self.img_channels, self.resolution // 2, self.resolution // 2])
|
||||
# img = upfirdn2d.upsample2d(img, self.resample_filter)
|
||||
if self.is_last or self.architecture == 'skip':
|
||||
y = self.torgb(x, next(w_iter), fused_modconv=fused_modconv)
|
||||
y = y.to(
|
||||
dtype=torch.float32, memory_format=torch.contiguous_format)
|
||||
img = img.add_(y) if img is not None else y
|
||||
|
||||
assert x.dtype == dtype
|
||||
assert img is None or img.dtype == torch.float32
|
||||
return x, img
|
||||
|
||||
def extra_repr(self):
|
||||
return f'resolution={self.resolution:d}, architecture={self.architecture:s}'
|
||||
|
||||
|
||||
# for 512x512 generation
|
||||
@persistence.persistent_class
|
||||
class SuperresolutionHybrid8XDC(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
channels,
|
||||
img_resolution,
|
||||
sr_num_fp16_res,
|
||||
sr_antialias,
|
||||
num_fp16_res=4,
|
||||
conv_clamp=None,
|
||||
channel_base=None,
|
||||
channel_max=None, # IGNORE
|
||||
**block_kwargs):
|
||||
super().__init__()
|
||||
assert img_resolution == 512
|
||||
|
||||
use_fp16 = sr_num_fp16_res > 0
|
||||
self.input_resolution = 128
|
||||
self.sr_antialias = sr_antialias
|
||||
self.block0 = SynthesisBlock(
|
||||
channels,
|
||||
256,
|
||||
w_dim=512,
|
||||
resolution=256,
|
||||
img_channels=3,
|
||||
is_last=False,
|
||||
use_fp16=use_fp16,
|
||||
conv_clamp=(256 if use_fp16 else None),
|
||||
**block_kwargs)
|
||||
self.block1 = SynthesisBlock(
|
||||
256,
|
||||
128,
|
||||
w_dim=512,
|
||||
resolution=512,
|
||||
img_channels=3,
|
||||
is_last=True,
|
||||
use_fp16=use_fp16,
|
||||
conv_clamp=(256 if use_fp16 else None),
|
||||
**block_kwargs)
|
||||
|
||||
def forward(self, rgb, x, ws, **block_kwargs):
|
||||
ws = ws[:, -1:, :].repeat(1, 3, 1)
|
||||
|
||||
if x.shape[-1] != self.input_resolution:
|
||||
x = torch.nn.functional.interpolate(
|
||||
x,
|
||||
size=(self.input_resolution, self.input_resolution),
|
||||
mode='bilinear',
|
||||
align_corners=False,
|
||||
antialias=self.sr_antialias)
|
||||
rgb = torch.nn.functional.interpolate(
|
||||
rgb,
|
||||
size=(self.input_resolution, self.input_resolution),
|
||||
mode='bilinear',
|
||||
align_corners=False,
|
||||
antialias=self.sr_antialias)
|
||||
|
||||
x, rgb = self.block0(x, rgb, ws, **block_kwargs)
|
||||
x, rgb = self.block1(x, rgb, ws, **block_kwargs)
|
||||
return rgb
|
||||
@@ -0,0 +1,242 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
|
||||
import torch
|
||||
|
||||
from modelscope.ops.image_control_3d_portrait.torch_utils import persistence
|
||||
from .networks_stylegan2 import FullyConnectedLayer
|
||||
from .networks_stylegan2 import Generator as StyleGAN2Backbone
|
||||
from .superresolution import SuperresolutionHybrid8XDC
|
||||
from .volumetric_rendering.ray_sampler import RaySampler
|
||||
from .volumetric_rendering.renderer import ImportanceRenderer
|
||||
|
||||
|
||||
@persistence.persistent_class
|
||||
class TriPlaneGenerator(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
z_dim, # Input latent (Z) dimensionality.
|
||||
c_dim, # Conditioning label (C) dimensionality.
|
||||
w_dim, # Intermediate latent (W) dimensionality.
|
||||
img_resolution, # Output resolution.
|
||||
img_channels, # Number of output color channels.
|
||||
sr_num_fp16_res=0,
|
||||
mapping_kwargs={}, # Arguments for MappingNetwork.
|
||||
rendering_kwargs={},
|
||||
sr_kwargs={},
|
||||
**synthesis_kwargs, # Arguments for SynthesisNetwork.
|
||||
):
|
||||
super().__init__()
|
||||
self.z_dim = z_dim
|
||||
self.c_dim = c_dim
|
||||
self.w_dim = w_dim
|
||||
self.img_resolution = img_resolution
|
||||
self.img_channels = img_channels
|
||||
self.renderer = ImportanceRenderer()
|
||||
self.ray_sampler = RaySampler()
|
||||
self.backbone = StyleGAN2Backbone(
|
||||
z_dim,
|
||||
c_dim,
|
||||
w_dim,
|
||||
img_resolution=256,
|
||||
img_channels=32 * 3,
|
||||
mapping_kwargs=mapping_kwargs,
|
||||
**synthesis_kwargs)
|
||||
self.superresolution = SuperresolutionHybrid8XDC(
|
||||
channels=32,
|
||||
img_resolution=img_resolution,
|
||||
sr_num_fp16_res=sr_num_fp16_res,
|
||||
sr_antialias=rendering_kwargs['sr_antialias'],
|
||||
**sr_kwargs)
|
||||
self.decoder = OSGDecoder(
|
||||
32, {
|
||||
'decoder_lr_mul': rendering_kwargs.get('decoder_lr_mul', 1),
|
||||
'decoder_output_dim': 32
|
||||
})
|
||||
self.neural_rendering_resolution = 64
|
||||
self.rendering_kwargs = rendering_kwargs
|
||||
|
||||
self._last_planes = None
|
||||
|
||||
def mapping(self,
|
||||
z,
|
||||
c,
|
||||
truncation_psi=1,
|
||||
truncation_cutoff=None,
|
||||
update_emas=False):
|
||||
if self.rendering_kwargs['c_gen_conditioning_zero']:
|
||||
c = torch.zeros_like(c)
|
||||
return self.backbone.mapping(
|
||||
z,
|
||||
c * self.rendering_kwargs.get('c_scale', 0),
|
||||
truncation_psi=truncation_psi,
|
||||
truncation_cutoff=truncation_cutoff,
|
||||
update_emas=update_emas)
|
||||
|
||||
def synthesis(self,
|
||||
ws,
|
||||
c,
|
||||
neural_rendering_resolution=None,
|
||||
update_emas=False,
|
||||
cache_backbone=False,
|
||||
use_cached_backbone=False,
|
||||
**synthesis_kwargs):
|
||||
cam2world_matrix = c[:, :16].view(-1, 4, 4)
|
||||
intrinsics = c[:, 16:25].view(-1, 3, 3)
|
||||
|
||||
if neural_rendering_resolution is None:
|
||||
neural_rendering_resolution = self.neural_rendering_resolution
|
||||
else:
|
||||
self.neural_rendering_resolution = neural_rendering_resolution
|
||||
|
||||
# Create a batch of rays for volume rendering
|
||||
ray_origins, ray_directions = self.ray_sampler(
|
||||
cam2world_matrix, intrinsics, neural_rendering_resolution)
|
||||
|
||||
# Create triplanes by running StyleGAN backbone
|
||||
N, M, _ = ray_origins.shape
|
||||
if use_cached_backbone and self._last_planes is not None:
|
||||
planes = self._last_planes
|
||||
else:
|
||||
planes = self.backbone.synthesis(
|
||||
ws, update_emas=update_emas, **synthesis_kwargs)
|
||||
if cache_backbone:
|
||||
self._last_planes = planes
|
||||
|
||||
# Reshape output into three 32-channel planes
|
||||
planes = planes.view(
|
||||
len(planes), 3, 32, planes.shape[-2], planes.shape[-1])
|
||||
|
||||
# Perform volume rendering
|
||||
feature_samples, depth_samples, weights_samples = self.renderer(
|
||||
planes, self.decoder, ray_origins, ray_directions,
|
||||
self.rendering_kwargs) # channels last
|
||||
|
||||
# Reshape into 'raw' neural-rendered image
|
||||
H = W = self.neural_rendering_resolution
|
||||
feature_image = feature_samples.permute(0, 2, 1).reshape(
|
||||
N, feature_samples.shape[-1], H, W).contiguous()
|
||||
depth_image = depth_samples.permute(0, 2, 1).reshape(N, 1, H, W)
|
||||
|
||||
# Run superresolution to get final image
|
||||
rgb_image = feature_image[:, :3]
|
||||
sr_image = self.superresolution(
|
||||
rgb_image,
|
||||
feature_image,
|
||||
ws,
|
||||
noise_mode=self.rendering_kwargs['superresolution_noise_mode'],
|
||||
**{
|
||||
k: synthesis_kwargs[k]
|
||||
for k in synthesis_kwargs.keys() if k != 'noise_mode'
|
||||
})
|
||||
|
||||
return {
|
||||
'image': sr_image,
|
||||
'image_raw': rgb_image,
|
||||
'image_depth': depth_image
|
||||
}
|
||||
|
||||
def sample(self,
|
||||
coordinates,
|
||||
directions,
|
||||
z,
|
||||
c,
|
||||
truncation_psi=1,
|
||||
truncation_cutoff=None,
|
||||
update_emas=False,
|
||||
**synthesis_kwargs):
|
||||
# Compute RGB features, density for arbitrary 3D coordinates. Mostly used for extracting shapes.
|
||||
ws = self.mapping(
|
||||
z,
|
||||
c,
|
||||
truncation_psi=truncation_psi,
|
||||
truncation_cutoff=truncation_cutoff,
|
||||
update_emas=update_emas)
|
||||
planes = self.backbone.synthesis(
|
||||
ws, update_emas=update_emas, **synthesis_kwargs)
|
||||
planes = planes.view(
|
||||
len(planes), 3, 32, planes.shape[-2], planes.shape[-1])
|
||||
return self.renderer.run_model(planes, self.decoder, coordinates,
|
||||
directions, self.rendering_kwargs)
|
||||
|
||||
def sample_mixed(self,
|
||||
coordinates,
|
||||
directions,
|
||||
ws,
|
||||
truncation_psi=1,
|
||||
truncation_cutoff=None,
|
||||
update_emas=False,
|
||||
**synthesis_kwargs):
|
||||
# Same as sample, but expects latent vectors 'ws' instead of Gaussian noise 'z'
|
||||
planes = self.backbone.synthesis(
|
||||
ws, update_emas=update_emas, **synthesis_kwargs)
|
||||
planes = planes.view(
|
||||
len(planes), 3, 32, planes.shape[-2], planes.shape[-1])
|
||||
return self.renderer.run_model(planes, self.decoder, coordinates,
|
||||
directions, self.rendering_kwargs)
|
||||
|
||||
def forward(self,
|
||||
z,
|
||||
c,
|
||||
truncation_psi=1,
|
||||
truncation_cutoff=None,
|
||||
neural_rendering_resolution=None,
|
||||
update_emas=False,
|
||||
cache_backbone=False,
|
||||
use_cached_backbone=False,
|
||||
**synthesis_kwargs):
|
||||
# Render a batch of generated images.
|
||||
ws = self.mapping(
|
||||
z,
|
||||
c,
|
||||
truncation_psi=truncation_psi,
|
||||
truncation_cutoff=truncation_cutoff,
|
||||
update_emas=update_emas)
|
||||
return self.synthesis(
|
||||
ws,
|
||||
c,
|
||||
update_emas=update_emas,
|
||||
neural_rendering_resolution=neural_rendering_resolution,
|
||||
cache_backbone=cache_backbone,
|
||||
use_cached_backbone=use_cached_backbone,
|
||||
**synthesis_kwargs)
|
||||
|
||||
|
||||
class OSGDecoder(torch.nn.Module):
|
||||
|
||||
def __init__(self, n_features, options):
|
||||
super().__init__()
|
||||
self.hidden_dim = 64
|
||||
|
||||
self.net = torch.nn.Sequential(
|
||||
FullyConnectedLayer(
|
||||
n_features,
|
||||
self.hidden_dim,
|
||||
lr_multiplier=options['decoder_lr_mul']), torch.nn.Softplus(),
|
||||
FullyConnectedLayer(
|
||||
self.hidden_dim,
|
||||
1 + options['decoder_output_dim'],
|
||||
lr_multiplier=options['decoder_lr_mul']))
|
||||
|
||||
def forward(self, sampled_features, ray_directions):
|
||||
# Aggregate features
|
||||
sampled_features = sampled_features.mean(1)
|
||||
x = sampled_features
|
||||
|
||||
N, M, C = x.shape
|
||||
x = x.view(N * M, C)
|
||||
|
||||
x = self.net(x)
|
||||
x = x.view(N, M, -1)
|
||||
rgb = torch.sigmoid(x[..., 1:]) * (
|
||||
1 + 2 * 0.001) - 0.001 # Uses sigmoid clamping from MipNeRF
|
||||
sigma = x[..., 0:1]
|
||||
return {'rgb': rgb, 'sigma': sigma}
|
||||
@@ -0,0 +1,697 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
from functools import partial
|
||||
|
||||
import segmentation_models_pytorch as smp
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from timm.models.layers import DropPath, to_2tuple, trunc_normal_
|
||||
|
||||
from .networks_stylegan2 import FullyConnectedLayer
|
||||
from .superresolution import SuperresolutionHybrid8XDC
|
||||
from .volumetric_rendering.ray_sampler import RaySampler
|
||||
from .volumetric_rendering.renderer import ImportanceRenderer
|
||||
|
||||
|
||||
class DWConv(nn.Module):
|
||||
|
||||
def __init__(self, dim=768):
|
||||
super(DWConv, self).__init__()
|
||||
self.dwconv = nn.Conv2d(dim, dim, 3, 1, 1, bias=True, groups=dim)
|
||||
|
||||
def forward(self, x, H, W):
|
||||
B, N, C = x.shape
|
||||
x = x.transpose(1, 2).view(B, C, H, W)
|
||||
x = self.dwconv(x)
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class Mlp(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
in_features,
|
||||
hidden_features=None,
|
||||
out_features=None,
|
||||
act_layer=nn.GELU,
|
||||
drop=0.):
|
||||
super().__init__()
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
self.fc1 = nn.Linear(in_features, hidden_features)
|
||||
self.dwconv = DWConv(hidden_features)
|
||||
self.act = act_layer()
|
||||
self.fc2 = nn.Linear(hidden_features, out_features)
|
||||
self.drop = nn.Dropout(drop)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x, H, W):
|
||||
x = self.fc1(x)
|
||||
x = self.dwconv(x, H, W)
|
||||
x = self.act(x)
|
||||
x = self.drop(x)
|
||||
x = self.fc2(x)
|
||||
x = self.drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
num_heads=8,
|
||||
qkv_bias=False,
|
||||
qk_scale=None,
|
||||
attn_drop=0.,
|
||||
proj_drop=0.,
|
||||
sr_ratio=1):
|
||||
super().__init__()
|
||||
assert dim % num_heads == 0, f'dim {dim} should be divided by num_heads {num_heads}.'
|
||||
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
self.scale = qk_scale or head_dim**-0.5
|
||||
|
||||
self.q = nn.Linear(dim, dim, bias=qkv_bias)
|
||||
self.kv = nn.Linear(dim, dim * 2, bias=qkv_bias)
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
|
||||
self.sr_ratio = sr_ratio
|
||||
if sr_ratio > 1:
|
||||
self.sr = nn.Conv2d(
|
||||
dim, dim, kernel_size=sr_ratio, stride=sr_ratio)
|
||||
self.norm = nn.LayerNorm(dim)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x, H, W):
|
||||
B, N, C = x.shape
|
||||
q = self.q(x).reshape(B, N, self.num_heads,
|
||||
C // self.num_heads).permute(0, 2, 1, 3)
|
||||
|
||||
if self.sr_ratio > 1:
|
||||
x_ = x.permute(0, 2, 1).reshape(B, C, H, W)
|
||||
x_ = self.sr(x_).reshape(B, C, -1).permute(0, 2, 1)
|
||||
x_ = self.norm(x_)
|
||||
kv = self.kv(x_).reshape(B, -1, 2, self.num_heads,
|
||||
C // self.num_heads).permute(
|
||||
2, 0, 3, 1, 4)
|
||||
else:
|
||||
kv = self.kv(x).reshape(B, -1, 2, self.num_heads,
|
||||
C // self.num_heads).permute(
|
||||
2, 0, 3, 1, 4)
|
||||
k, v = kv[0], kv[1]
|
||||
|
||||
attn = (q @ k.transpose(-2, -1)) * self.scale
|
||||
attn = attn.softmax(dim=-1)
|
||||
attn = self.attn_drop(attn)
|
||||
|
||||
x = (attn @ v).transpose(1, 2).reshape(B, N, C)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class Block(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
num_heads,
|
||||
mlp_ratio=4.,
|
||||
qkv_bias=False,
|
||||
qk_scale=None,
|
||||
drop=0.,
|
||||
attn_drop=0.,
|
||||
drop_path=0.,
|
||||
act_layer=nn.GELU,
|
||||
norm_layer=nn.LayerNorm,
|
||||
sr_ratio=1):
|
||||
super().__init__()
|
||||
self.norm1 = norm_layer(dim)
|
||||
self.attn = Attention(
|
||||
dim,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
attn_drop=attn_drop,
|
||||
proj_drop=drop,
|
||||
sr_ratio=sr_ratio)
|
||||
self.drop_path = DropPath(
|
||||
drop_path) if drop_path > 0. else nn.Identity()
|
||||
self.norm2 = norm_layer(dim)
|
||||
mlp_hidden_dim = int(dim * mlp_ratio)
|
||||
self.mlp = Mlp(
|
||||
in_features=dim,
|
||||
hidden_features=mlp_hidden_dim,
|
||||
act_layer=act_layer,
|
||||
drop=drop)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x, H, W):
|
||||
x = x + self.drop_path(self.attn(self.norm1(x), H, W))
|
||||
x = x + self.drop_path(self.mlp(self.norm2(x), H, W))
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class OverlapPatchEmbed(nn.Module):
|
||||
""" Image to Patch Embedding
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
img_size=224,
|
||||
patch_size=7,
|
||||
stride=4,
|
||||
in_chans=3,
|
||||
embed_dim=768):
|
||||
super().__init__()
|
||||
img_size = to_2tuple(img_size)
|
||||
patch_size = to_2tuple(patch_size)
|
||||
|
||||
self.img_size = img_size
|
||||
self.patch_size = patch_size
|
||||
self.H, self.W = img_size[0] // patch_size[0], img_size[
|
||||
1] // patch_size[1]
|
||||
self.num_patches = self.H * self.W
|
||||
self.proj = nn.Conv2d(
|
||||
in_chans,
|
||||
embed_dim,
|
||||
kernel_size=patch_size,
|
||||
stride=stride,
|
||||
padding=(patch_size[0] // 2, patch_size[1] // 2))
|
||||
self.norm = nn.LayerNorm(embed_dim)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.proj(x)
|
||||
_, _, H, W = x.shape
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
x = self.norm(x)
|
||||
|
||||
return x, H, W
|
||||
|
||||
|
||||
class Encoder_low(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
img_size=64,
|
||||
depth=5,
|
||||
in_chans=256,
|
||||
embed_dims=1024,
|
||||
num_head=4,
|
||||
mlp_ratio=2,
|
||||
sr_ratio=1,
|
||||
qkv_bias=True,
|
||||
qk_scale=None,
|
||||
drop_rate=0.,
|
||||
attn_drop_rate=0.,
|
||||
drop_path_rate=0.,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6)):
|
||||
super().__init__()
|
||||
self.depth = depth
|
||||
|
||||
self.deeplabnet = smp.DeepLabV3(
|
||||
encoder_name='resnet34',
|
||||
encoder_depth=5,
|
||||
encoder_weights=None,
|
||||
decoder_channels=256,
|
||||
in_channels=5,
|
||||
classes=1)
|
||||
|
||||
self.deeplabnet.encoder.conv1 = nn.Conv2d(
|
||||
5,
|
||||
64,
|
||||
kernel_size=(7, 7),
|
||||
stride=(2, 2),
|
||||
padding=(3, 3),
|
||||
bias=False)
|
||||
self.deeplabnet.segmentation_head = nn.Sequential()
|
||||
self.deeplabnet.encoder.bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer1[0].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer1[0].bn2 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer1[1].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer1[1].bn2 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer1[2].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer1[2].bn2 = nn.Sequential()
|
||||
|
||||
self.deeplabnet.encoder.layer2[0].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer2[0].bn2 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer2[0].downsample[1] = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer2[1].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer2[1].bn2 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer2[2].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer2[2].bn2 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer2[3].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer2[3].bn2 = nn.Sequential()
|
||||
|
||||
self.deeplabnet.encoder.layer3[0].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer3[0].bn2 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer3[0].downsample[1] = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer3[1].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer3[1].bn2 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer3[2].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer3[2].bn2 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer3[3].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer3[3].bn2 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer3[4].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer3[4].bn2 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer3[5].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer3[5].bn2 = nn.Sequential()
|
||||
|
||||
self.deeplabnet.encoder.layer4[0].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer4[0].bn2 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer4[0].downsample[1] = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer4[1].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer4[1].bn2 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer4[2].bn1 = nn.Sequential()
|
||||
self.deeplabnet.encoder.layer4[2].bn2 = nn.Sequential()
|
||||
|
||||
self.deeplabnet.decoder[0].convs[0][1] = nn.Sequential()
|
||||
self.deeplabnet.decoder[0].convs[1][1] = nn.Sequential()
|
||||
self.deeplabnet.decoder[0].convs[2][1] = nn.Sequential()
|
||||
self.deeplabnet.decoder[0].convs[3][1] = nn.Sequential()
|
||||
self.deeplabnet.decoder[0].convs[4][2] = nn.Sequential()
|
||||
self.deeplabnet.decoder[0].project[1] = nn.Sequential()
|
||||
self.deeplabnet.decoder[2] = nn.Sequential()
|
||||
|
||||
self.patch_embed = OverlapPatchEmbed(
|
||||
img_size=img_size,
|
||||
patch_size=3,
|
||||
stride=2,
|
||||
in_chans=in_chans,
|
||||
embed_dim=embed_dims)
|
||||
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)]
|
||||
cur = 0
|
||||
self.vit_block = nn.ModuleList([
|
||||
Block(
|
||||
dim=embed_dims,
|
||||
num_heads=num_head,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop_rate,
|
||||
attn_drop=attn_drop_rate,
|
||||
drop_path=dpr[cur + i],
|
||||
norm_layer=norm_layer,
|
||||
sr_ratio=sr_ratio) for i in range(depth)
|
||||
])
|
||||
self.norm1 = norm_layer(embed_dims)
|
||||
self.ps = nn.PixelShuffle(2)
|
||||
|
||||
self.upsample1 = nn.UpsamplingBilinear2d(scale_factor=2)
|
||||
self.conv1 = nn.Conv2d(256, 128, kernel_size=3, stride=1, padding=1)
|
||||
self.relu1 = nn.ReLU()
|
||||
self.upsample2 = nn.UpsamplingBilinear2d(scale_factor=2)
|
||||
self.conv2 = nn.Conv2d(128, 128, kernel_size=3, stride=1, padding=1)
|
||||
self.relu2 = nn.ReLU()
|
||||
self.conv3 = nn.Conv2d(128, 96, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, input):
|
||||
B = input.shape[0]
|
||||
|
||||
f_low = self.deeplabnet(input)
|
||||
x, H, W = self.patch_embed(f_low)
|
||||
|
||||
for i, blk in enumerate(self.vit_block):
|
||||
x = blk(x, H, W)
|
||||
x = self.norm1(x)
|
||||
x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
|
||||
x = self.ps(x)
|
||||
|
||||
x = self.relu1(self.conv1(self.upsample1(x)))
|
||||
x = self.relu2(self.conv2(self.upsample2(x)))
|
||||
x = self.conv3(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class Encoder_high(nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.conv1 = nn.Conv2d(5, 64, kernel_size=7, stride=2, padding=3)
|
||||
self.relu1 = nn.LeakyReLU(0.01)
|
||||
self.conv2 = nn.Conv2d(64, 96, kernel_size=3, stride=1, padding=1)
|
||||
self.relu2 = nn.LeakyReLU(0.01)
|
||||
self.conv3 = nn.Conv2d(96, 96, kernel_size=3, stride=1, padding=1)
|
||||
self.relu3 = nn.LeakyReLU(0.01)
|
||||
self.conv4 = nn.Conv2d(96, 96, kernel_size=3, stride=1, padding=1)
|
||||
self.relu4 = nn.LeakyReLU(0.01)
|
||||
self.conv5 = nn.Conv2d(96, 96, kernel_size=3, stride=1, padding=1)
|
||||
self.relu5 = nn.LeakyReLU(0.01)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.relu1(self.conv1(x))
|
||||
x = self.relu2(self.conv2(x))
|
||||
x = self.relu3(self.conv3(x))
|
||||
x = self.relu4(self.conv4(x))
|
||||
x = self.relu5(self.conv5(x))
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class MixFeature(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
img_size=256,
|
||||
depth=1,
|
||||
in_chans=128,
|
||||
embed_dims=1024,
|
||||
num_head=2,
|
||||
mlp_ratio=2,
|
||||
sr_ratio=2,
|
||||
qkv_bias=True,
|
||||
qk_scale=None,
|
||||
drop_rate=0.,
|
||||
attn_drop_rate=0.,
|
||||
drop_path_rate=0.,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6)):
|
||||
super().__init__()
|
||||
self.conv1 = nn.Conv2d(192, 256, kernel_size=3, stride=1, padding=1)
|
||||
self.relu1 = nn.LeakyReLU(0.01)
|
||||
self.conv2 = nn.Conv2d(256, 128, kernel_size=3, stride=1, padding=1)
|
||||
self.relu2 = nn.LeakyReLU(0.01)
|
||||
|
||||
self.patch_embed = OverlapPatchEmbed(
|
||||
img_size=img_size,
|
||||
patch_size=3,
|
||||
stride=2,
|
||||
in_chans=in_chans,
|
||||
embed_dim=embed_dims)
|
||||
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)]
|
||||
cur = 0
|
||||
self.vit_block = nn.ModuleList([
|
||||
Block(
|
||||
dim=embed_dims,
|
||||
num_heads=num_head,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop_rate,
|
||||
attn_drop=attn_drop_rate,
|
||||
drop_path=dpr[cur + i],
|
||||
norm_layer=norm_layer,
|
||||
sr_ratio=sr_ratio) for i in range(depth)
|
||||
])
|
||||
self.norm1 = norm_layer(embed_dims)
|
||||
self.ps = nn.PixelShuffle(2)
|
||||
|
||||
self.conv3 = nn.Conv2d(352, 256, kernel_size=3, stride=1, padding=1)
|
||||
self.relu3 = nn.LeakyReLU(0.01)
|
||||
self.conv4 = nn.Conv2d(256, 128, kernel_size=3, stride=1, padding=1)
|
||||
self.relu4 = nn.LeakyReLU(0.01)
|
||||
self.conv5 = nn.Conv2d(128, 128, kernel_size=3, stride=1, padding=1)
|
||||
self.relu5 = nn.LeakyReLU(0.01)
|
||||
self.conv6 = nn.Conv2d(128, 96, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x_low, x_high):
|
||||
x = torch.cat((x_low, x_high), 1)
|
||||
B = x.shape[0]
|
||||
|
||||
x = self.relu1(self.conv1(x))
|
||||
x = self.relu2(self.conv2(x))
|
||||
|
||||
x, H, W = self.patch_embed(x)
|
||||
|
||||
for i, blk in enumerate(self.vit_block):
|
||||
x = blk(x, H, W)
|
||||
x = self.norm1(x)
|
||||
x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
|
||||
x = self.ps(x)
|
||||
|
||||
x = torch.cat((x, x_low), 1)
|
||||
x = self.relu3(self.conv3(x))
|
||||
x = self.relu4(self.conv4(x))
|
||||
x = self.relu5(self.conv5(x))
|
||||
x = self.conv6(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class OSGDecoder(torch.nn.Module):
|
||||
|
||||
def __init__(self, n_features, options):
|
||||
super().__init__()
|
||||
self.hidden_dim = 64
|
||||
|
||||
self.net = torch.nn.Sequential(
|
||||
FullyConnectedLayer(
|
||||
n_features,
|
||||
self.hidden_dim,
|
||||
lr_multiplier=options['decoder_lr_mul']), torch.nn.Softplus(),
|
||||
FullyConnectedLayer(
|
||||
self.hidden_dim,
|
||||
1 + options['decoder_output_dim'],
|
||||
lr_multiplier=options['decoder_lr_mul']))
|
||||
|
||||
def forward(self, sampled_features, ray_directions):
|
||||
sampled_features = sampled_features.mean(1)
|
||||
x = sampled_features
|
||||
|
||||
N, M, C = x.shape
|
||||
x = x.view(N * M, C)
|
||||
|
||||
x = self.net(x)
|
||||
x = x.view(N, M, -1)
|
||||
rgb = torch.sigmoid(x[..., 1:]) * (1 + 2 * 0.001) - 0.001
|
||||
sigma = x[..., 0:1]
|
||||
return {'rgb': rgb, 'sigma': sigma}
|
||||
|
||||
|
||||
class TriplaneEncoder(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
img_resolution,
|
||||
sr_num_fp16_res=0,
|
||||
rendering_kwargs={},
|
||||
sr_kwargs={}):
|
||||
super().__init__()
|
||||
self.encoder_low = Encoder_low(
|
||||
img_size=64,
|
||||
depth=5,
|
||||
in_chans=256,
|
||||
embed_dims=1024,
|
||||
num_head=4,
|
||||
mlp_ratio=2,
|
||||
sr_ratio=1,
|
||||
qkv_bias=True,
|
||||
qk_scale=None,
|
||||
drop_rate=0.,
|
||||
attn_drop_rate=0.,
|
||||
drop_path_rate=0.,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6))
|
||||
self.encoder_high = Encoder_high()
|
||||
self.mix = MixFeature(
|
||||
img_size=256,
|
||||
depth=1,
|
||||
in_chans=128,
|
||||
embed_dims=1024,
|
||||
num_head=2,
|
||||
mlp_ratio=2,
|
||||
sr_ratio=2,
|
||||
qkv_bias=True,
|
||||
qk_scale=None,
|
||||
drop_rate=0.,
|
||||
attn_drop_rate=0.,
|
||||
drop_path_rate=0.,
|
||||
norm_layer=partial(nn.LayerNorm, eps=1e-6))
|
||||
|
||||
self.renderer = ImportanceRenderer()
|
||||
self.ray_sampler = RaySampler()
|
||||
self.superresolution = SuperresolutionHybrid8XDC(
|
||||
channels=32,
|
||||
img_resolution=img_resolution,
|
||||
sr_num_fp16_res=sr_num_fp16_res,
|
||||
sr_antialias=rendering_kwargs['sr_antialias'],
|
||||
**sr_kwargs)
|
||||
self.decoder = OSGDecoder(
|
||||
32, {
|
||||
'decoder_lr_mul': rendering_kwargs.get('decoder_lr_mul', 1),
|
||||
'decoder_output_dim': 32
|
||||
})
|
||||
self.neural_rendering_resolution = 128
|
||||
self.rendering_kwargs = rendering_kwargs
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def gen_interfeats(self, ws, planes, camera_params):
|
||||
planes = planes.view(
|
||||
len(planes), 3, 32, planes.shape[-2], planes.shape[-1])
|
||||
|
||||
cam2world_matrix = camera_params[:, :16].view(-1, 4, 4)
|
||||
intrinsics = camera_params[:, 16:25].view(-1, 3, 3)
|
||||
H = W = self.neural_rendering_resolution
|
||||
ray_origins, ray_directions = self.ray_sampler(
|
||||
cam2world_matrix, intrinsics, self.neural_rendering_resolution)
|
||||
N, M, _ = ray_origins.shape
|
||||
feature_samples, depth_samples, weights_samples = self.renderer(
|
||||
planes, self.decoder, ray_origins, ray_directions,
|
||||
self.rendering_kwargs)
|
||||
feature_image = feature_samples.permute(0, 2, 1).reshape(
|
||||
N, feature_samples.shape[-1], H, W).contiguous()
|
||||
depth_image = depth_samples.permute(0, 2, 1).reshape(N, 1, H, W)
|
||||
|
||||
rgb_image = feature_image[:, :3]
|
||||
sr_image = self.superresolution(
|
||||
rgb_image, feature_image, ws, noise_mode='const')
|
||||
|
||||
return depth_image, feature_image, rgb_image, sr_image
|
||||
|
||||
def sample(self, coordinates, directions, planes):
|
||||
planes = planes.view(
|
||||
len(planes), 3, 32, planes.shape[-2], planes.shape[-1])
|
||||
return self.renderer.run_model(planes, self.decoder, coordinates,
|
||||
directions, self.rendering_kwargs)
|
||||
|
||||
def forward(self, ws, x, camera_ref, camera_mv):
|
||||
f = self.encoder_low(x)
|
||||
f_high = self.encoder_high(x)
|
||||
planes = self.mix(f, f_high)
|
||||
|
||||
depth_ref, feature_ref, rgb_ref, sr_ref = self.gen_interfeats(
|
||||
ws, planes, camera_ref)
|
||||
if camera_mv is not None:
|
||||
depth_mv, feature_mv, rgb_mv, sr_mv = self.gen_interfeats(
|
||||
ws, planes, camera_mv)
|
||||
else:
|
||||
depth_mv = feature_mv = rgb_mv = sr_mv = None
|
||||
|
||||
return planes, depth_ref, feature_ref, rgb_ref, sr_ref, depth_mv, feature_mv, rgb_mv, sr_mv
|
||||
|
||||
|
||||
def get_parameter_number(net):
|
||||
total_num = sum(p.numel() for p in net.parameters())
|
||||
trainable_num = sum(p.numel() for p in net.parameters() if p.requires_grad)
|
||||
return {'Total': total_num, 'Trainable': trainable_num}
|
||||
@@ -0,0 +1,11 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
|
||||
# empty
|
||||
@@ -0,0 +1,137 @@
|
||||
# MIT License
|
||||
|
||||
# Copyright (c) 2022 Petr Kellnhofer
|
||||
|
||||
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
# of this software and associated documentation files (the "Software"), to deal
|
||||
# in the Software without restriction, including without limitation the rights
|
||||
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
# copies of the Software, and to permit persons to whom the Software is
|
||||
# furnished to do so, subject to the following conditions:
|
||||
|
||||
# The above copyright notice and this permission notice shall be included in all
|
||||
# copies or substantial portions of the Software.
|
||||
|
||||
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
# SOFTWARE.
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def transform_vectors(matrix: torch.Tensor,
|
||||
vectors4: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Left-multiplies MxM @ NxM. Returns NxM.
|
||||
"""
|
||||
res = torch.matmul(vectors4, matrix.T)
|
||||
return res
|
||||
|
||||
|
||||
def normalize_vecs(vectors: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Normalize vector lengths.
|
||||
"""
|
||||
return vectors / (torch.norm(vectors, dim=-1, keepdim=True))
|
||||
|
||||
|
||||
def torch_dot(x: torch.Tensor, y: torch.Tensor):
|
||||
"""
|
||||
Dot product of two tensors.
|
||||
"""
|
||||
return (x * y).sum(-1)
|
||||
|
||||
|
||||
def get_ray_limits_box(rays_o: torch.Tensor, rays_d: torch.Tensor,
|
||||
box_side_length):
|
||||
"""
|
||||
Author: Petr Kellnhofer
|
||||
Intersects rays with the [-1, 1] NDC volume.
|
||||
Returns min and max distance of entry.
|
||||
Returns -1 for no intersection.
|
||||
https://www.scratchapixel.com/lessons/3d-basic-rendering/minimal-ray-tracer-rendering-simple-shapes/ray-box-intersection
|
||||
"""
|
||||
o_shape = rays_o.shape
|
||||
rays_o = rays_o.detach().reshape(-1, 3)
|
||||
rays_d = rays_d.detach().reshape(-1, 3)
|
||||
|
||||
temp_min_1 = -1 * (box_side_length / 2)
|
||||
temp_min_2 = -1 * (box_side_length / 2)
|
||||
temp_min_3 = -1 * (box_side_length / 2)
|
||||
bb_min = [temp_min_1, temp_min_2, temp_min_3]
|
||||
temp_max_1 = 1 * (box_side_length / 2)
|
||||
temp_max_2 = 1 * (box_side_length / 2)
|
||||
temp_max_3 = 1 * (box_side_length / 2)
|
||||
bb_max = [temp_max_1, temp_max_2, temp_max_3]
|
||||
bounds = torch.tensor([bb_min, bb_max],
|
||||
dtype=rays_o.dtype,
|
||||
device=rays_o.device)
|
||||
is_valid = torch.ones(rays_o.shape[:-1], dtype=bool, device=rays_o.device)
|
||||
|
||||
# Precompute inverse for stability.
|
||||
invdir = 1 / rays_d
|
||||
sign = (invdir < 0).long()
|
||||
|
||||
# Intersect with YZ plane.
|
||||
tmin = (bounds.index_select(0, sign[..., 0])[..., 0]
|
||||
- rays_o[..., 0]) * invdir[..., 0]
|
||||
tmax = (bounds.index_select(0, 1 - sign[..., 0])[..., 0]
|
||||
- rays_o[..., 0]) * invdir[..., 0]
|
||||
|
||||
# Intersect with XZ plane.
|
||||
tymin = (bounds.index_select(0, sign[..., 1])[..., 1]
|
||||
- rays_o[..., 1]) * invdir[..., 1]
|
||||
tymax = (bounds.index_select(0, 1 - sign[..., 1])[..., 1]
|
||||
- rays_o[..., 1]) * invdir[..., 1]
|
||||
|
||||
# Resolve parallel rays.
|
||||
is_valid[torch.logical_or(tmin > tymax, tymin > tmax)] = False
|
||||
|
||||
# Use the shortest intersection.
|
||||
tmin = torch.max(tmin, tymin)
|
||||
tmax = torch.min(tmax, tymax)
|
||||
|
||||
# Intersect with XY plane.
|
||||
tzmin = (bounds.index_select(0, sign[..., 2])[..., 2]
|
||||
- rays_o[..., 2]) * invdir[..., 2]
|
||||
tzmax = (bounds.index_select(0, 1 - sign[..., 2])[..., 2]
|
||||
- rays_o[..., 2]) * invdir[..., 2]
|
||||
|
||||
# Resolve parallel rays.
|
||||
is_valid[torch.logical_or(tmin > tzmax, tzmin > tmax)] = False
|
||||
|
||||
# Use the shortest intersection.
|
||||
tmin = torch.max(tmin, tzmin)
|
||||
tmax = torch.min(tmax, tzmax)
|
||||
|
||||
# Mark invalid.
|
||||
tmin[torch.logical_not(is_valid)] = -1
|
||||
tmax[torch.logical_not(is_valid)] = -2
|
||||
|
||||
return tmin.reshape(*o_shape[:-1], 1), tmax.reshape(*o_shape[:-1], 1)
|
||||
|
||||
|
||||
def linspace(start: torch.Tensor, stop: torch.Tensor, num: int):
|
||||
"""
|
||||
Creates a tensor of shape [num, *start.shape] whose values are evenly spaced from start to end, inclusive.
|
||||
Replicates but the multi-dimensional bahaviour of numpy.linspace in PyTorch.
|
||||
"""
|
||||
# create a tensor of 'num' steps from 0 to 1
|
||||
steps = torch.arange(
|
||||
num, dtype=torch.float32, device=start.device) / (
|
||||
num - 1)
|
||||
|
||||
# reshape the 'steps' tensor to [-1, *([1]*start.ndim)] to allow for broadcastings
|
||||
# - using 'steps.reshape([-1, *([1]*start.ndim)])' would be nice here but torchscript
|
||||
# "cannot statically infer the expected size of a list in this contex", hence the code below
|
||||
for i in range(start.ndim):
|
||||
steps = steps.unsqueeze(-1)
|
||||
|
||||
# the output starts at 'start' and increments until 'stop' in each dimension
|
||||
out = start[None] + steps * (stop - start)[None]
|
||||
|
||||
return out
|
||||
@@ -0,0 +1,67 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""
|
||||
The ray marcher takes the raw output of the implicit representation and
|
||||
uses the volume rendering equation to produce composited colors and depths.
|
||||
Based off of the implementation in MipNeRF (this one doesn't do any cone tracing though!)
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class MipRayMarcher2(nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def run_forward(self, colors, densities, depths, rendering_options):
|
||||
deltas = depths[:, :, 1:] - depths[:, :, :-1]
|
||||
colors_mid = (colors[:, :, :-1] + colors[:, :, 1:]) / 2
|
||||
densities_mid = (densities[:, :, :-1] + densities[:, :, 1:]) / 2
|
||||
depths_mid = (depths[:, :, :-1] + depths[:, :, 1:]) / 2
|
||||
|
||||
if rendering_options['clamp_mode'] == 'softplus':
|
||||
densities_mid = F.softplus(
|
||||
densities_mid
|
||||
- 1) # activation bias of -1 makes things initialize better
|
||||
else:
|
||||
assert False, 'MipRayMarcher only supports `clamp_mode`=`softplus`!'
|
||||
|
||||
density_delta = densities_mid * deltas
|
||||
|
||||
alpha = 1 - torch.exp(-density_delta)
|
||||
|
||||
alpha_shifted = torch.cat(
|
||||
[torch.ones_like(alpha[:, :, :1]), 1 - alpha + 1e-10], -2)
|
||||
weights = alpha * torch.cumprod(alpha_shifted, -2)[:, :, :-1]
|
||||
|
||||
composite_rgb = torch.sum(weights * colors_mid, -2)
|
||||
weight_total = weights.sum(2)
|
||||
composite_depth = torch.sum(weights * depths_mid, -2) / weight_total
|
||||
|
||||
# clip the composite to min/max range of depths
|
||||
composite_depth = torch.nan_to_num(composite_depth, float('inf'))
|
||||
composite_depth = torch.clamp(composite_depth, torch.min(depths),
|
||||
torch.max(depths))
|
||||
|
||||
if rendering_options.get('white_back', False):
|
||||
composite_rgb = composite_rgb + 1 - weight_total
|
||||
|
||||
composite_rgb = composite_rgb * 2 - 1 # Scale to (-1, 1)
|
||||
|
||||
return composite_rgb, composite_depth, weights
|
||||
|
||||
def forward(self, colors, densities, depths, rendering_options):
|
||||
composite_rgb, composite_depth, weights = self.run_forward(
|
||||
colors, densities, depths, rendering_options)
|
||||
|
||||
return composite_rgb, composite_depth, weights
|
||||
@@ -0,0 +1,80 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""
|
||||
The ray sampler is a module that takes in camera matrices and resolution and batches of rays.
|
||||
Expects cam2world matrices that use the OpenCV camera coordinate system conventions.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class RaySampler(torch.nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.ray_origins_h, self.ray_directions, self.depths, self.image_coords, self.rendering_options = \
|
||||
None, None, None, None, None
|
||||
|
||||
def forward(self, cam2world_matrix, intrinsics, resolution):
|
||||
"""
|
||||
Create batches of rays and return origins and directions.
|
||||
|
||||
cam2world_matrix: (N, 4, 4)
|
||||
intrinsics: (N, 3, 3)
|
||||
resolution: int
|
||||
|
||||
ray_origins: (N, M, 3)
|
||||
ray_dirs: (N, M, 2)
|
||||
"""
|
||||
N, M = cam2world_matrix.shape[0], resolution**2
|
||||
cam_locs_world = cam2world_matrix[:, :3, 3]
|
||||
fx = intrinsics[:, 0, 0]
|
||||
fy = intrinsics[:, 1, 1]
|
||||
cx = intrinsics[:, 0, 2]
|
||||
cy = intrinsics[:, 1, 2]
|
||||
sk = intrinsics[:, 0, 1]
|
||||
|
||||
uv = torch.stack(
|
||||
torch.meshgrid(
|
||||
torch.arange(
|
||||
resolution,
|
||||
dtype=torch.float32,
|
||||
device=cam2world_matrix.device),
|
||||
torch.arange(
|
||||
resolution,
|
||||
dtype=torch.float32,
|
||||
device=cam2world_matrix.device))) * (1. / resolution) + (
|
||||
0.5 / resolution)
|
||||
uv = uv.flip(0).reshape(2, -1).transpose(1, 0)
|
||||
uv = uv.unsqueeze(0).repeat(cam2world_matrix.shape[0], 1, 1)
|
||||
|
||||
x_cam = uv[:, :, 0].view(N, -1)
|
||||
y_cam = uv[:, :, 1].view(N, -1)
|
||||
z_cam = torch.ones((N, M), device=cam2world_matrix.device)
|
||||
|
||||
x_lift = (x_cam - cx.unsqueeze(-1) + cy.unsqueeze(-1)
|
||||
* sk.unsqueeze(-1) / fy.unsqueeze(-1) - sk.unsqueeze(-1)
|
||||
* y_cam / fy.unsqueeze(-1)) / fx.unsqueeze(-1) * z_cam
|
||||
y_lift = (y_cam - cy.unsqueeze(-1)) / fy.unsqueeze(-1) * z_cam
|
||||
|
||||
cam_rel_points = torch.stack(
|
||||
(x_lift, y_lift, z_cam, torch.ones_like(z_cam)), dim=-1)
|
||||
|
||||
world_rel_points = torch.bmm(cam2world_matrix,
|
||||
cam_rel_points.permute(0, 2, 1)).permute(
|
||||
0, 2, 1)[:, :, :3]
|
||||
|
||||
ray_dirs = world_rel_points - cam_locs_world[:, None, :]
|
||||
ray_dirs = torch.nn.functional.normalize(ray_dirs, dim=2)
|
||||
|
||||
ray_origins = cam_locs_world.unsqueeze(1).repeat(
|
||||
1, ray_dirs.shape[1], 1)
|
||||
|
||||
return ray_origins, ray_dirs
|
||||
@@ -0,0 +1,341 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""
|
||||
The renderer is a module that takes in rays, decides where to sample along each
|
||||
ray, and computes pixel colors using the volume rendering equation.
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from . import math_utils
|
||||
from .ray_marcher import MipRayMarcher2
|
||||
|
||||
|
||||
def generate_planes():
|
||||
"""
|
||||
Defines planes by the three vectors that form the "axes" of the
|
||||
plane. Should work with arbitrary number of planes and planes of
|
||||
arbitrary orientation.
|
||||
"""
|
||||
return torch.tensor(
|
||||
[[[1, 0, 0], [0, 1, 0], [0, 0, 1]], [[1, 0, 0], [0, 0, 1], [0, 1, 0]],
|
||||
[[0, 0, 1], [1, 0, 0], [0, 1, 0]]],
|
||||
dtype=torch.float32)
|
||||
|
||||
|
||||
def project_onto_planes(planes, coordinates):
|
||||
"""
|
||||
Does a projection of a 3D point onto a batch of 2D planes,
|
||||
returning 2D plane coordinates.
|
||||
|
||||
Takes plane axes of shape n_planes, 3, 3
|
||||
# Takes coordinates of shape N, M, 3
|
||||
# returns projections of shape N*n_planes, M, 2
|
||||
"""
|
||||
N, M, C = coordinates.shape
|
||||
n_planes, _, _ = planes.shape
|
||||
coordinates = coordinates.unsqueeze(1).expand(-1, n_planes, -1,
|
||||
-1).reshape(
|
||||
N * n_planes, M, 3)
|
||||
inv_planes = torch.linalg.inv(planes).unsqueeze(0).expand(
|
||||
N, -1, -1, -1).reshape(N * n_planes, 3, 3).to(coordinates.device)
|
||||
projections = torch.bmm(coordinates, inv_planes)
|
||||
return projections[..., :2]
|
||||
|
||||
|
||||
def sample_from_planes(plane_axes,
|
||||
plane_features,
|
||||
coordinates,
|
||||
mode='bilinear',
|
||||
padding_mode='zeros',
|
||||
box_warp=None):
|
||||
assert padding_mode == 'zeros'
|
||||
N, n_planes, C, H, W = plane_features.shape
|
||||
_, M, _ = coordinates.shape
|
||||
plane_features = plane_features.view(N * n_planes, C, H, W)
|
||||
|
||||
coordinates = (2 / box_warp) * coordinates # TODO: add specific box bounds
|
||||
|
||||
projected_coordinates = project_onto_planes(plane_axes,
|
||||
coordinates).unsqueeze(1)
|
||||
output_features = torch.nn.functional.grid_sample(
|
||||
plane_features,
|
||||
projected_coordinates.float(),
|
||||
mode=mode,
|
||||
padding_mode=padding_mode,
|
||||
align_corners=False).permute(0, 3, 2, 1).reshape(N, n_planes, M, C)
|
||||
return output_features
|
||||
|
||||
|
||||
def sample_from_3dgrid(grid, coordinates):
|
||||
"""
|
||||
Expects coordinates in shape (batch_size, num_points_per_batch, 3)
|
||||
Expects grid in shape (1, channels, H, W, D)
|
||||
(Also works if grid has batch size)
|
||||
Returns sampled features of shape (batch_size, num_points_per_batch, feature_channels)
|
||||
"""
|
||||
batch_size, n_coords, n_dims = coordinates.shape
|
||||
sampled_features = torch.nn.functional.grid_sample(
|
||||
grid.expand(batch_size, -1, -1, -1, -1),
|
||||
coordinates.reshape(batch_size, 1, 1, -1, n_dims),
|
||||
mode='bilinear',
|
||||
padding_mode='zeros',
|
||||
align_corners=False)
|
||||
N, C, H, W, D = sampled_features.shape
|
||||
sampled_features = sampled_features.permute(0, 4, 3, 2,
|
||||
1).reshape(N, H * W * D, C)
|
||||
return sampled_features
|
||||
|
||||
|
||||
class ImportanceRenderer(torch.nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.ray_marcher = MipRayMarcher2()
|
||||
self.plane_axes = generate_planes()
|
||||
|
||||
def forward(self, planes, decoder, ray_origins, ray_directions,
|
||||
rendering_options):
|
||||
self.plane_axes = self.plane_axes.to(ray_origins.device)
|
||||
|
||||
if rendering_options['ray_start'] == rendering_options[
|
||||
'ray_end'] == 'auto':
|
||||
ray_start, ray_end = math_utils.get_ray_limits_box(
|
||||
ray_origins,
|
||||
ray_directions,
|
||||
box_side_length=rendering_options['box_warp'])
|
||||
is_ray_valid = ray_end > ray_start
|
||||
if torch.any(is_ray_valid).item():
|
||||
ray_start[~is_ray_valid] = ray_start[is_ray_valid].min()
|
||||
ray_end[~is_ray_valid] = ray_start[is_ray_valid].max()
|
||||
depths_coarse = self.sample_stratified(
|
||||
ray_origins, ray_start, ray_end,
|
||||
rendering_options['depth_resolution'],
|
||||
rendering_options['disparity_space_sampling'])
|
||||
else:
|
||||
# Create stratified depth samples
|
||||
depths_coarse = self.sample_stratified(
|
||||
ray_origins, rendering_options['ray_start'],
|
||||
rendering_options['ray_end'],
|
||||
rendering_options['depth_resolution'],
|
||||
rendering_options['disparity_space_sampling'])
|
||||
|
||||
batch_size, num_rays, samples_per_ray, _ = depths_coarse.shape
|
||||
|
||||
# Coarse Pass
|
||||
sample_coordinates = (
|
||||
ray_origins.unsqueeze(-2)
|
||||
+ depths_coarse * ray_directions.unsqueeze(-2)).reshape(
|
||||
batch_size, -1, 3)
|
||||
sample_directions = ray_directions.unsqueeze(-2).expand(
|
||||
-1, -1, samples_per_ray, -1).reshape(batch_size, -1, 3)
|
||||
|
||||
out = self.run_model(planes, decoder, sample_coordinates,
|
||||
sample_directions, rendering_options)
|
||||
colors_coarse = out['rgb']
|
||||
densities_coarse = out['sigma']
|
||||
colors_coarse = colors_coarse.reshape(batch_size, num_rays,
|
||||
samples_per_ray,
|
||||
colors_coarse.shape[-1])
|
||||
densities_coarse = densities_coarse.reshape(batch_size, num_rays,
|
||||
samples_per_ray, 1)
|
||||
|
||||
# Fine Pass
|
||||
N_importance = rendering_options['depth_resolution_importance']
|
||||
if N_importance > 0:
|
||||
_, _, weights = self.ray_marcher(colors_coarse, densities_coarse,
|
||||
depths_coarse, rendering_options)
|
||||
|
||||
depths_fine = self.sample_importance(depths_coarse, weights,
|
||||
N_importance)
|
||||
|
||||
sample_directions = ray_directions.unsqueeze(-2).expand(
|
||||
-1, -1, N_importance, -1).reshape(batch_size, -1, 3)
|
||||
sample_coordinates = (
|
||||
ray_origins.unsqueeze(-2)
|
||||
+ depths_fine * ray_directions.unsqueeze(-2)).reshape(
|
||||
batch_size, -1, 3)
|
||||
|
||||
out = self.run_model(planes, decoder, sample_coordinates,
|
||||
sample_directions, rendering_options)
|
||||
colors_fine = out['rgb']
|
||||
densities_fine = out['sigma']
|
||||
colors_fine = colors_fine.reshape(batch_size, num_rays,
|
||||
N_importance,
|
||||
colors_fine.shape[-1])
|
||||
densities_fine = densities_fine.reshape(batch_size, num_rays,
|
||||
N_importance, 1)
|
||||
|
||||
all_depths, all_colors, all_densities = self.unify_samples(
|
||||
depths_coarse, colors_coarse, densities_coarse, depths_fine,
|
||||
colors_fine, densities_fine)
|
||||
|
||||
# Aggregate
|
||||
rgb_final, depth_final, weights = self.ray_marcher(
|
||||
all_colors, all_densities, all_depths, rendering_options)
|
||||
else:
|
||||
rgb_final, depth_final, weights = self.ray_marcher(
|
||||
colors_coarse, densities_coarse, depths_coarse,
|
||||
rendering_options)
|
||||
|
||||
return rgb_final, depth_final, weights.sum(2)
|
||||
|
||||
def run_model(self, planes, decoder, sample_coordinates, sample_directions,
|
||||
options):
|
||||
sampled_features = sample_from_planes(
|
||||
self.plane_axes,
|
||||
planes,
|
||||
sample_coordinates,
|
||||
padding_mode='zeros',
|
||||
box_warp=options['box_warp'])
|
||||
|
||||
out = decoder(sampled_features, sample_directions)
|
||||
if options.get('density_noise', 0) > 0:
|
||||
out['sigma'] += torch.randn_like(
|
||||
out['sigma']) * options['density_noise']
|
||||
return out
|
||||
|
||||
def sort_samples(self, all_depths, all_colors, all_densities):
|
||||
_, indices = torch.sort(all_depths, dim=-2)
|
||||
all_depths = torch.gather(all_depths, -2, indices)
|
||||
all_colors = torch.gather(
|
||||
all_colors, -2, indices.expand(-1, -1, -1, all_colors.shape[-1]))
|
||||
all_densities = torch.gather(all_densities, -2,
|
||||
indices.expand(-1, -1, -1, 1))
|
||||
return all_depths, all_colors, all_densities
|
||||
|
||||
def unify_samples(self, depths1, colors1, densities1, depths2, colors2,
|
||||
densities2):
|
||||
all_depths = torch.cat([depths1, depths2], dim=-2)
|
||||
all_colors = torch.cat([colors1, colors2], dim=-2)
|
||||
all_densities = torch.cat([densities1, densities2], dim=-2)
|
||||
|
||||
_, indices = torch.sort(all_depths, dim=-2)
|
||||
all_depths = torch.gather(all_depths, -2, indices)
|
||||
all_colors = torch.gather(
|
||||
all_colors, -2, indices.expand(-1, -1, -1, all_colors.shape[-1]))
|
||||
all_densities = torch.gather(all_densities, -2,
|
||||
indices.expand(-1, -1, -1, 1))
|
||||
|
||||
return all_depths, all_colors, all_densities
|
||||
|
||||
def sample_stratified(self,
|
||||
ray_origins,
|
||||
ray_start,
|
||||
ray_end,
|
||||
depth_resolution,
|
||||
disparity_space_sampling=False):
|
||||
"""
|
||||
Return depths of approximately uniformly spaced samples along rays.
|
||||
"""
|
||||
N, M, _ = ray_origins.shape
|
||||
if disparity_space_sampling:
|
||||
depths_coarse = torch.linspace(
|
||||
0, 1, depth_resolution,
|
||||
device=ray_origins.device).reshape(1, 1, depth_resolution,
|
||||
1).repeat(N, M, 1, 1)
|
||||
depth_delta = 1 / (depth_resolution - 1)
|
||||
depths_coarse += torch.rand_like(depths_coarse) * depth_delta
|
||||
depths_coarse = 1. / (1. / ray_start * (1. - depths_coarse)
|
||||
+ 1. / ray_end * depths_coarse)
|
||||
else:
|
||||
if type(ray_start) == torch.Tensor:
|
||||
depths_coarse = math_utils.linspace(ray_start, ray_end,
|
||||
depth_resolution).permute(
|
||||
1, 2, 0, 3)
|
||||
depth_delta = (ray_end - ray_start) / (depth_resolution - 1)
|
||||
depths_coarse += torch.rand_like(depths_coarse) * depth_delta[
|
||||
..., None]
|
||||
else:
|
||||
depths_coarse = torch.linspace(
|
||||
ray_start,
|
||||
ray_end,
|
||||
depth_resolution,
|
||||
device=ray_origins.device).reshape(1, 1, depth_resolution,
|
||||
1).repeat(N, M, 1, 1)
|
||||
depth_delta = (ray_end - ray_start) / (depth_resolution - 1)
|
||||
depths_coarse += torch.rand_like(depths_coarse) * depth_delta
|
||||
|
||||
return depths_coarse
|
||||
|
||||
def sample_importance(self, z_vals, weights, N_importance):
|
||||
"""
|
||||
Return depths of importance sampled points along rays. See NeRF importance sampling for more.
|
||||
"""
|
||||
with torch.no_grad():
|
||||
batch_size, num_rays, samples_per_ray, _ = z_vals.shape
|
||||
|
||||
z_vals = z_vals.reshape(batch_size * num_rays, samples_per_ray)
|
||||
weights = weights.reshape(
|
||||
batch_size * num_rays,
|
||||
-1) # -1 to account for loss of 1 sample in MipRayMarcher
|
||||
|
||||
# smooth weights
|
||||
weights = torch.nn.functional.max_pool1d(
|
||||
weights.unsqueeze(1).float(), 2, 1, padding=1)
|
||||
weights = torch.nn.functional.avg_pool1d(weights, 2, 1).squeeze()
|
||||
weights = weights + 0.01
|
||||
|
||||
z_vals_mid = 0.5 * (z_vals[:, :-1] + z_vals[:, 1:])
|
||||
importance_z_vals = self.sample_pdf(z_vals_mid, weights[:, 1:-1],
|
||||
N_importance).detach().reshape(
|
||||
batch_size, num_rays,
|
||||
N_importance, 1)
|
||||
return importance_z_vals
|
||||
|
||||
def sample_pdf(self, bins, weights, N_importance, det=False, eps=1e-5):
|
||||
"""
|
||||
Sample @N_importance samples from @bins with distribution defined by @weights.
|
||||
Inputs:
|
||||
bins: (N_rays, N_samples_+1) where N_samples_ is "the number of coarse samples per ray - 2"
|
||||
weights: (N_rays, N_samples_)
|
||||
N_importance: the number of samples to draw from the distribution
|
||||
det: deterministic or not
|
||||
eps: a small number to prevent division by zero
|
||||
Outputs:
|
||||
samples: the sampled samples
|
||||
"""
|
||||
N_rays, N_samples_ = weights.shape
|
||||
weights = weights + eps # prevent division by zero (don't do inplace op!)
|
||||
pdf = weights / torch.sum(
|
||||
weights, -1, keepdim=True) # (N_rays, N_samples_)
|
||||
cdf = torch.cumsum(
|
||||
pdf, -1) # (N_rays, N_samples), cumulative distribution function
|
||||
cdf = torch.cat([torch.zeros_like(cdf[:, :1]), cdf],
|
||||
-1) # (N_rays, N_samples_+1)
|
||||
# padded to 0~1 inclusive
|
||||
|
||||
if det:
|
||||
u = torch.linspace(0, 1, N_importance, device=bins.device)
|
||||
u = u.expand(N_rays, N_importance)
|
||||
else:
|
||||
u = torch.rand(N_rays, N_importance, device=bins.device)
|
||||
u = u.contiguous()
|
||||
|
||||
inds = torch.searchsorted(cdf, u, right=True)
|
||||
below = torch.clamp_min(inds - 1, 0)
|
||||
above = torch.clamp_max(inds, N_samples_)
|
||||
|
||||
inds_sampled = torch.stack([below, above],
|
||||
-1).view(N_rays, 2 * N_importance)
|
||||
cdf_g = torch.gather(cdf, 1,
|
||||
inds_sampled).view(N_rays, N_importance, 2)
|
||||
bins_g = torch.gather(bins, 1,
|
||||
inds_sampled).view(N_rays, N_importance, 2)
|
||||
|
||||
denom = cdf_g[..., 1] - cdf_g[..., 0]
|
||||
denom[denom < eps] = 1
|
||||
|
||||
samples = bins_g[..., 0] + (u - cdf_g[..., 0]) / denom * (
|
||||
bins_g[..., 1] - bins_g[..., 0])
|
||||
return samples
|
||||
11
modelscope/ops/image_control_3d_portrait/dnnlib/__init__.py
Normal file
11
modelscope/ops/image_control_3d_portrait/dnnlib/__init__.py
Normal file
@@ -0,0 +1,11 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
|
||||
from .util import EasyDict, make_cache_dir_path
|
||||
52
modelscope/ops/image_control_3d_portrait/dnnlib/util.py
Normal file
52
modelscope/ops/image_control_3d_portrait/dnnlib/util.py
Normal file
@@ -0,0 +1,52 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""Miscellaneous utility classes and functions."""
|
||||
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
from typing import Any, List, Tuple, Union
|
||||
|
||||
|
||||
class EasyDict(dict):
|
||||
"""Convenience class that behaves like a dict but allows access with the attribute syntax."""
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
try:
|
||||
return self[name]
|
||||
except KeyError:
|
||||
raise AttributeError(name)
|
||||
|
||||
def __setattr__(self, name: str, value: Any) -> None:
|
||||
self[name] = value
|
||||
|
||||
def __delattr__(self, name: str) -> None:
|
||||
del self[name]
|
||||
|
||||
|
||||
_dnnlib_cache_dir = None
|
||||
|
||||
|
||||
def set_cache_dir(path: str) -> None:
|
||||
global _dnnlib_cache_dir
|
||||
_dnnlib_cache_dir = path
|
||||
|
||||
|
||||
def make_cache_dir_path(*paths: str) -> str:
|
||||
if _dnnlib_cache_dir is not None:
|
||||
return os.path.join(_dnnlib_cache_dir, *paths)
|
||||
if 'DNNLIB_CACHE_DIR' in os.environ:
|
||||
return os.path.join(os.environ['DNNLIB_CACHE_DIR'], *paths)
|
||||
if 'HOME' in os.environ:
|
||||
return os.path.join(os.environ['HOME'], '.cache', 'dnnlib', *paths)
|
||||
if 'USERPROFILE' in os.environ:
|
||||
return os.path.join(os.environ['USERPROFILE'], '.cache', 'dnnlib',
|
||||
*paths)
|
||||
return os.path.join(tempfile.gettempdir(), '.cache', 'dnnlib', *paths)
|
||||
@@ -0,0 +1,181 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
|
||||
import glob
|
||||
import hashlib
|
||||
import importlib
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import uuid
|
||||
|
||||
import torch
|
||||
import torch.utils.cpp_extension
|
||||
from torch.utils.file_baton import FileBaton
|
||||
|
||||
# Global options.
|
||||
|
||||
verbosity = 'brief' # Verbosity level: 'none', 'brief', 'full'
|
||||
|
||||
# Internal helper funcs.
|
||||
|
||||
|
||||
def _find_compiler_bindir():
|
||||
patterns = [
|
||||
'C:/Program Files (x86)/Microsoft Visual Studio/*/Professional/VC/Tools/MSVC/*/bin/Hostx64/x64',
|
||||
'C:/Program Files (x86)/Microsoft Visual Studio/*/BuildTools/VC/Tools/MSVC/*/bin/Hostx64/x64',
|
||||
'C:/Program Files (x86)/Microsoft Visual Studio/*/Community/VC/Tools/MSVC/*/bin/Hostx64/x64',
|
||||
'C:/Program Files (x86)/Microsoft Visual Studio */vc/bin',
|
||||
]
|
||||
for pattern in patterns:
|
||||
matches = sorted(glob.glob(pattern))
|
||||
if len(matches):
|
||||
return matches[-1]
|
||||
return None
|
||||
|
||||
|
||||
def _get_mangled_gpu_name():
|
||||
name = torch.cuda.get_device_name().lower()
|
||||
out = []
|
||||
for c in name:
|
||||
if re.match('[a-z0-9_-]+', c):
|
||||
out.append(c)
|
||||
else:
|
||||
out.append('-')
|
||||
return ''.join(out)
|
||||
|
||||
|
||||
# Main entry point for compiling and loading C++/CUDA plugins.
|
||||
|
||||
_cached_plugins = dict()
|
||||
|
||||
|
||||
def get_plugin(module_name,
|
||||
sources,
|
||||
headers=None,
|
||||
source_dir=None,
|
||||
**build_kwargs):
|
||||
assert verbosity in ['none', 'brief', 'full']
|
||||
if headers is None:
|
||||
headers = []
|
||||
if source_dir is not None:
|
||||
sources = [os.path.join(source_dir, fname) for fname in sources]
|
||||
headers = [os.path.join(source_dir, fname) for fname in headers]
|
||||
|
||||
# Already cached?
|
||||
if module_name in _cached_plugins:
|
||||
return _cached_plugins[module_name]
|
||||
|
||||
# Print status.
|
||||
if verbosity == 'full':
|
||||
print(f'Setting up PyTorch plugin "{module_name}"...')
|
||||
elif verbosity == 'brief':
|
||||
print(
|
||||
f'Setting up PyTorch plugin "{module_name}"... ',
|
||||
end='',
|
||||
flush=True)
|
||||
verbose_build = (verbosity == 'full')
|
||||
|
||||
# Compile and load.
|
||||
try:
|
||||
if os.name == 'nt' and os.system('where cl.exe >nul 2>nul') != 0:
|
||||
compiler_bindir = _find_compiler_bindir()
|
||||
if compiler_bindir is None:
|
||||
raise RuntimeError(
|
||||
f'Could not find MSVC/GCC/CLANG installation on this computer.'
|
||||
f' Check _find_compiler_bindir() in "{__file__}".')
|
||||
os.environ['PATH'] += ';' + compiler_bindir
|
||||
|
||||
# Some containers set TORCH_CUDA_ARCH_LIST to a list that can either
|
||||
# break the build or unnecessarily restrict what's available to nvcc.
|
||||
# Unset it to let nvcc decide based on what's available on the
|
||||
# machine.
|
||||
os.environ['TORCH_CUDA_ARCH_LIST'] = ''
|
||||
|
||||
# Incremental build md5sum trickery. Copies all the input source files
|
||||
# into a cached build directory under a combined md5 digest of the input
|
||||
# source files. Copying is done only if the combined digest has changed.
|
||||
# This keeps input file timestamps and filenames the same as in previous
|
||||
# extension builds, allowing for fast incremental rebuilds.
|
||||
#
|
||||
# This optimization is done only in case all the source files reside in
|
||||
# a single directory (just for simplicity) and if the TORCH_EXTENSIONS_DIR
|
||||
# environment variable is set (we take this as a signal that the user
|
||||
# actually cares about this.)
|
||||
#
|
||||
# EDIT: We now do it regardless of TORCH_EXTENSIOS_DIR, in order to work
|
||||
# around the *.cu dependency bug in ninja config.
|
||||
#
|
||||
all_source_files = sorted(sources + headers)
|
||||
all_source_dirs = set(
|
||||
os.path.dirname(fname) for fname in all_source_files)
|
||||
if len(all_source_dirs
|
||||
) == 1: # and ('TORCH_EXTENSIONS_DIR' in os.environ):
|
||||
|
||||
# Compute combined hash digest for all source files.
|
||||
hash_md5 = hashlib.md5()
|
||||
for src in all_source_files:
|
||||
with open(src, 'rb') as f:
|
||||
hash_md5.update(f.read())
|
||||
|
||||
# Select cached build directory name.
|
||||
source_digest = hash_md5.hexdigest()
|
||||
build_top_dir = torch.utils.cpp_extension._get_build_directory(
|
||||
module_name, verbose=verbose_build) # pylint: disable=protected-access
|
||||
cached_build_dir = os.path.join(
|
||||
build_top_dir, f'{source_digest}-{_get_mangled_gpu_name()}')
|
||||
|
||||
if not os.path.isdir(cached_build_dir):
|
||||
tmpdir = f'{build_top_dir}/srctmp-{uuid.uuid4().hex}'
|
||||
os.makedirs(tmpdir)
|
||||
for src in all_source_files:
|
||||
shutil.copyfile(
|
||||
src, os.path.join(tmpdir, os.path.basename(src)))
|
||||
try:
|
||||
os.replace(tmpdir, cached_build_dir) # atomic
|
||||
except OSError:
|
||||
# source directory already exists, delete tmpdir and its contents.
|
||||
shutil.rmtree(tmpdir)
|
||||
if not os.path.isdir(cached_build_dir):
|
||||
raise
|
||||
|
||||
# Compile.
|
||||
cached_sources = [
|
||||
os.path.join(cached_build_dir, os.path.basename(fname))
|
||||
for fname in sources
|
||||
]
|
||||
torch.utils.cpp_extension.load(
|
||||
name=module_name,
|
||||
build_directory=cached_build_dir,
|
||||
verbose=verbose_build,
|
||||
sources=cached_sources,
|
||||
**build_kwargs)
|
||||
else:
|
||||
torch.utils.cpp_extension.load(
|
||||
name=module_name,
|
||||
verbose=verbose_build,
|
||||
sources=sources,
|
||||
**build_kwargs)
|
||||
|
||||
# Load.
|
||||
module = importlib.import_module(module_name)
|
||||
|
||||
except Exception:
|
||||
if verbosity == 'brief':
|
||||
print('Failed!')
|
||||
raise
|
||||
|
||||
# Print status and add to cache dict.
|
||||
if verbosity == 'full':
|
||||
print(f'Done setting up PyTorch plugin "{module_name}".')
|
||||
elif verbosity == 'brief':
|
||||
print('Done.')
|
||||
_cached_plugins[module_name] = module
|
||||
return module
|
||||
325
modelscope/ops/image_control_3d_portrait/torch_utils/misc.py
Normal file
325
modelscope/ops/image_control_3d_portrait/torch_utils/misc.py
Normal file
@@ -0,0 +1,325 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
|
||||
import contextlib
|
||||
import re
|
||||
import warnings
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from .. import dnnlib
|
||||
|
||||
# Cached construction of constant tensors. Avoids CPU=>GPU copy when the
|
||||
# same constant is used multiple times.
|
||||
|
||||
_constant_cache = dict()
|
||||
|
||||
|
||||
def constant(value, shape=None, dtype=None, device=None, memory_format=None):
|
||||
value = np.asarray(value)
|
||||
if shape is not None:
|
||||
shape = tuple(shape)
|
||||
if dtype is None:
|
||||
dtype = torch.get_default_dtype()
|
||||
if device is None:
|
||||
device = torch.device('cpu')
|
||||
if memory_format is None:
|
||||
memory_format = torch.contiguous_format
|
||||
|
||||
key = (value.shape, value.dtype, value.tobytes(), shape, dtype, device,
|
||||
memory_format)
|
||||
tensor = _constant_cache.get(key, None)
|
||||
if tensor is None:
|
||||
tensor = torch.as_tensor(value.copy(), dtype=dtype, device=device)
|
||||
if shape is not None:
|
||||
tensor, _ = torch.broadcast_tensors(tensor, torch.empty(shape))
|
||||
tensor = tensor.contiguous(memory_format=memory_format)
|
||||
_constant_cache[key] = tensor
|
||||
return tensor
|
||||
|
||||
|
||||
# Replace NaN/Inf with specified numerical values.
|
||||
|
||||
try:
|
||||
nan_to_num = torch.nan_to_num # 1.8.0a0
|
||||
except AttributeError:
|
||||
|
||||
def nan_to_num(input, nan=0.0, posinf=None, neginf=None, *, out=None): # pylint: disable=redefined-builtin
|
||||
assert isinstance(input, torch.Tensor)
|
||||
if posinf is None:
|
||||
posinf = torch.finfo(input.dtype).max
|
||||
if neginf is None:
|
||||
neginf = torch.finfo(input.dtype).min
|
||||
assert nan == 0
|
||||
return torch.clamp(
|
||||
input.unsqueeze(0).nansum(0), min=neginf, max=posinf, out=out)
|
||||
|
||||
|
||||
# Symbolic assert.
|
||||
|
||||
try:
|
||||
symbolic_assert = torch._assert # 1.8.0a0 # pylint: disable=protected-access
|
||||
except AttributeError:
|
||||
symbolic_assert = torch.Assert # 1.7.0
|
||||
|
||||
# Context manager to temporarily suppress known warnings in torch.jit.trace().
|
||||
# Note: Cannot use catch_warnings because of https://bugs.python.org/issue29672
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def suppress_tracer_warnings():
|
||||
flt = ('ignore', None, torch.jit.TracerWarning, None, 0)
|
||||
warnings.filters.insert(0, flt)
|
||||
yield
|
||||
warnings.filters.remove(flt)
|
||||
|
||||
|
||||
# Assert that the shape of a tensor matches the given list of integers.
|
||||
# None indicates that the size of a dimension is allowed to vary.
|
||||
# Performs symbolic assertion when used in torch.jit.trace().
|
||||
|
||||
|
||||
def assert_shape(tensor, ref_shape):
|
||||
if tensor.ndim != len(ref_shape):
|
||||
raise AssertionError(
|
||||
f'Wrong number of dimensions: got {tensor.ndim}, expected {len(ref_shape)}'
|
||||
)
|
||||
for idx, (size, ref_size) in enumerate(zip(tensor.shape, ref_shape)):
|
||||
if ref_size is None:
|
||||
pass
|
||||
elif isinstance(ref_size, torch.Tensor):
|
||||
with suppress_tracer_warnings(
|
||||
): # as_tensor results are registered as constants
|
||||
symbolic_assert(
|
||||
torch.equal(torch.as_tensor(size), ref_size),
|
||||
f'Wrong size for dimension {idx}')
|
||||
elif isinstance(size, torch.Tensor):
|
||||
with suppress_tracer_warnings(
|
||||
): # as_tensor results are registered as constants
|
||||
symbolic_assert(
|
||||
torch.equal(size, torch.as_tensor(ref_size)),
|
||||
f'Wrong size for dimension {idx}: expected {ref_size}')
|
||||
elif size != ref_size:
|
||||
raise AssertionError(
|
||||
f'Wrong size for dimension {idx}: got {size}, expected {ref_size}'
|
||||
)
|
||||
|
||||
|
||||
# Function decorator that calls torch.autograd.profiler.record_function().
|
||||
|
||||
|
||||
def profiled_function(fn):
|
||||
|
||||
def decorator(*args, **kwargs):
|
||||
with torch.autograd.profiler.record_function(fn.__name__):
|
||||
return fn(*args, **kwargs)
|
||||
|
||||
decorator.__name__ = fn.__name__
|
||||
return decorator
|
||||
|
||||
|
||||
# Sampler for torch.utils.data.DataLoader that loops over the dataset
|
||||
# indefinitely, shuffling items as it goes.
|
||||
|
||||
|
||||
class InfiniteSampler(torch.utils.data.Sampler):
|
||||
|
||||
def __init__(self,
|
||||
dataset,
|
||||
rank=0,
|
||||
num_replicas=1,
|
||||
shuffle=True,
|
||||
seed=0,
|
||||
window_size=0.5):
|
||||
assert len(dataset) > 0
|
||||
assert num_replicas > 0
|
||||
assert 0 <= rank < num_replicas
|
||||
assert 0 <= window_size <= 1
|
||||
super().__init__(dataset)
|
||||
self.dataset = dataset
|
||||
self.rank = rank
|
||||
self.num_replicas = num_replicas
|
||||
self.shuffle = shuffle
|
||||
self.seed = seed
|
||||
self.window_size = window_size
|
||||
|
||||
def __iter__(self):
|
||||
order = np.arange(len(self.dataset))
|
||||
rnd = None
|
||||
window = 0
|
||||
if self.shuffle:
|
||||
rnd = np.random.RandomState(self.seed)
|
||||
rnd.shuffle(order)
|
||||
window = int(np.rint(order.size * self.window_size))
|
||||
|
||||
idx = 0
|
||||
while True:
|
||||
i = idx % order.size
|
||||
if idx % self.num_replicas == self.rank:
|
||||
yield order[i]
|
||||
if window >= 2:
|
||||
j = (i - rnd.randint(window)) % order.size
|
||||
order[i], order[j] = order[j], order[i]
|
||||
idx += 1
|
||||
|
||||
|
||||
# Utilities for operating with torch.nn.Module parameters and buffers.
|
||||
|
||||
|
||||
def params_and_buffers(module):
|
||||
assert isinstance(module, torch.nn.Module)
|
||||
return list(module.parameters()) + list(module.buffers())
|
||||
|
||||
|
||||
def named_params_and_buffers(module):
|
||||
assert isinstance(module, torch.nn.Module)
|
||||
return list(module.named_parameters()) + list(module.named_buffers())
|
||||
|
||||
|
||||
def copy_params_and_buffers(src_module, dst_module, require_all=False):
|
||||
assert isinstance(src_module, torch.nn.Module)
|
||||
assert isinstance(dst_module, torch.nn.Module)
|
||||
src_tensors = dict(named_params_and_buffers(src_module))
|
||||
for name, tensor in named_params_and_buffers(dst_module):
|
||||
assert (name in src_tensors) or (not require_all)
|
||||
if name in src_tensors:
|
||||
tensor.copy_(src_tensors[name].detach()).requires_grad_(
|
||||
tensor.requires_grad)
|
||||
|
||||
|
||||
# Context manager for easily enabling/disabling DistributedDataParallel
|
||||
# synchronization.
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def ddp_sync(module, sync):
|
||||
assert isinstance(module, torch.nn.Module)
|
||||
if sync or not isinstance(module,
|
||||
torch.nn.parallel.DistributedDataParallel):
|
||||
yield
|
||||
else:
|
||||
with module.no_sync():
|
||||
yield
|
||||
|
||||
|
||||
# Check DistributedDataParallel consistency across processes.
|
||||
|
||||
|
||||
def check_ddp_consistency(module, ignore_regex=None):
|
||||
assert isinstance(module, torch.nn.Module)
|
||||
for name, tensor in named_params_and_buffers(module):
|
||||
fullname = type(module).__name__ + '.' + name
|
||||
if ignore_regex is not None and re.fullmatch(ignore_regex, fullname):
|
||||
continue
|
||||
tensor = tensor.detach()
|
||||
if tensor.is_floating_point():
|
||||
tensor = nan_to_num(tensor)
|
||||
other = tensor.clone()
|
||||
torch.distributed.broadcast(tensor=other, src=0)
|
||||
assert (tensor == other).all(), fullname
|
||||
|
||||
|
||||
# Print summary table of module hierarchy.
|
||||
|
||||
|
||||
def print_module_summary(module, inputs, max_nesting=3, skip_redundant=True):
|
||||
assert isinstance(module, torch.nn.Module)
|
||||
assert not isinstance(module, torch.jit.ScriptModule)
|
||||
assert isinstance(inputs, (tuple, list))
|
||||
|
||||
# Register hooks.
|
||||
entries = []
|
||||
nesting = [0]
|
||||
|
||||
def pre_hook(_mod, _inputs):
|
||||
nesting[0] += 1
|
||||
|
||||
def post_hook(mod, _inputs, outputs):
|
||||
nesting[0] -= 1
|
||||
if nesting[0] <= max_nesting:
|
||||
outputs = list(outputs) if isinstance(outputs,
|
||||
(tuple,
|
||||
list)) else [outputs]
|
||||
outputs = [t for t in outputs if isinstance(t, torch.Tensor)]
|
||||
entries.append(dnnlib.EasyDict(mod=mod, outputs=outputs))
|
||||
|
||||
hooks = [
|
||||
mod.register_forward_pre_hook(pre_hook) for mod in module.modules()
|
||||
]
|
||||
hooks += [mod.register_forward_hook(post_hook) for mod in module.modules()]
|
||||
|
||||
# Run module.
|
||||
outputs = module(*inputs)
|
||||
for hook in hooks:
|
||||
hook.remove()
|
||||
|
||||
# Identify unique outputs, parameters, and buffers.
|
||||
tensors_seen = set()
|
||||
for e in entries:
|
||||
e.unique_params = [
|
||||
t for t in e.mod.parameters() if id(t) not in tensors_seen
|
||||
]
|
||||
e.unique_buffers = [
|
||||
t for t in e.mod.buffers() if id(t) not in tensors_seen
|
||||
]
|
||||
e.unique_outputs = [t for t in e.outputs if id(t) not in tensors_seen]
|
||||
tensors_seen |= {
|
||||
id(t)
|
||||
for t in e.unique_params + e.unique_buffers + e.unique_outputs
|
||||
}
|
||||
|
||||
# Filter out redundant entries.
|
||||
if skip_redundant:
|
||||
entries = [
|
||||
e for e in entries if len(e.unique_params) or len(e.unique_buffers)
|
||||
or len(e.unique_outputs)
|
||||
]
|
||||
|
||||
# Construct table.
|
||||
rows = [[
|
||||
type(module).__name__, 'Parameters', 'Buffers', 'Output shape',
|
||||
'Datatype'
|
||||
]]
|
||||
rows += [['---'] * len(rows[0])]
|
||||
param_total = 0
|
||||
buffer_total = 0
|
||||
submodule_names = {mod: name for name, mod in module.named_modules()}
|
||||
for e in entries:
|
||||
name = '<top-level>' if e.mod is module else submodule_names[e.mod]
|
||||
param_size = sum(t.numel() for t in e.unique_params)
|
||||
buffer_size = sum(t.numel() for t in e.unique_buffers)
|
||||
output_shapes = [str(list(t.shape)) for t in e.outputs]
|
||||
output_dtypes = [str(t.dtype).split('.')[-1] for t in e.outputs]
|
||||
rows += [[
|
||||
name + (':0' if len(e.outputs) >= 2 else ''),
|
||||
str(param_size) if param_size else '-',
|
||||
str(buffer_size) if buffer_size else '-',
|
||||
(output_shapes + ['-'])[0],
|
||||
(output_dtypes + ['-'])[0],
|
||||
]]
|
||||
for idx in range(1, len(e.outputs)):
|
||||
rows += [[
|
||||
name + f':{idx}', '-', '-', output_shapes[idx],
|
||||
output_dtypes[idx]
|
||||
]]
|
||||
param_total += param_size
|
||||
buffer_total += buffer_size
|
||||
rows += [['---'] * len(rows[0])]
|
||||
rows += [['Total', str(param_total), str(buffer_total), '-', '-']]
|
||||
|
||||
# Print table.
|
||||
widths = [max(len(cell) for cell in column) for column in zip(*rows)]
|
||||
print()
|
||||
for row in rows:
|
||||
print(' '.join(cell + ' ' * (width - len(cell))
|
||||
for cell, width in zip(row, widths)))
|
||||
print()
|
||||
return outputs
|
||||
@@ -0,0 +1,103 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
*
|
||||
* NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
* property and proprietary rights in and to this material, related
|
||||
* documentation and any modifications thereto. Any use, reproduction,
|
||||
* disclosure or distribution of this material and related documentation
|
||||
* without an express license agreement from NVIDIA CORPORATION or
|
||||
* its affiliates is strictly prohibited.
|
||||
*/
|
||||
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include "bias_act.h"
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
|
||||
static bool has_same_layout(torch::Tensor x, torch::Tensor y)
|
||||
{
|
||||
if (x.dim() != y.dim())
|
||||
return false;
|
||||
for (int64_t i = 0; i < x.dim(); i++)
|
||||
{
|
||||
if (x.size(i) != y.size(i))
|
||||
return false;
|
||||
if (x.size(i) >= 2 && x.stride(i) != y.stride(i))
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
|
||||
static torch::Tensor bias_act(torch::Tensor x, torch::Tensor b, torch::Tensor xref, torch::Tensor yref, torch::Tensor dy, int grad, int dim, int act, float alpha, float gain, float clamp)
|
||||
{
|
||||
// Validate arguments.
|
||||
TORCH_CHECK(x.is_cuda(), "x must reside on CUDA device");
|
||||
TORCH_CHECK(b.numel() == 0 || (b.dtype() == x.dtype() && b.device() == x.device()), "b must have the same dtype and device as x");
|
||||
TORCH_CHECK(xref.numel() == 0 || (xref.sizes() == x.sizes() && xref.dtype() == x.dtype() && xref.device() == x.device()), "xref must have the same shape, dtype, and device as x");
|
||||
TORCH_CHECK(yref.numel() == 0 || (yref.sizes() == x.sizes() && yref.dtype() == x.dtype() && yref.device() == x.device()), "yref must have the same shape, dtype, and device as x");
|
||||
TORCH_CHECK(dy.numel() == 0 || (dy.sizes() == x.sizes() && dy.dtype() == x.dtype() && dy.device() == x.device()), "dy must have the same dtype and device as x");
|
||||
TORCH_CHECK(x.numel() <= INT_MAX, "x is too large");
|
||||
TORCH_CHECK(b.dim() == 1, "b must have rank 1");
|
||||
TORCH_CHECK(b.numel() == 0 || (dim >= 0 && dim < x.dim()), "dim is out of bounds");
|
||||
TORCH_CHECK(b.numel() == 0 || b.numel() == x.size(dim), "b has wrong number of elements");
|
||||
TORCH_CHECK(grad >= 0, "grad must be non-negative");
|
||||
|
||||
// Validate layout.
|
||||
TORCH_CHECK(x.is_non_overlapping_and_dense(), "x must be non-overlapping and dense");
|
||||
TORCH_CHECK(b.is_contiguous(), "b must be contiguous");
|
||||
TORCH_CHECK(xref.numel() == 0 || has_same_layout(xref, x), "xref must have the same layout as x");
|
||||
TORCH_CHECK(yref.numel() == 0 || has_same_layout(yref, x), "yref must have the same layout as x");
|
||||
TORCH_CHECK(dy.numel() == 0 || has_same_layout(dy, x), "dy must have the same layout as x");
|
||||
|
||||
// Create output tensor.
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(x));
|
||||
torch::Tensor y = torch::empty_like(x);
|
||||
TORCH_CHECK(has_same_layout(y, x), "y must have the same layout as x");
|
||||
|
||||
// Initialize CUDA kernel parameters.
|
||||
bias_act_kernel_params p;
|
||||
p.x = x.data_ptr();
|
||||
p.b = (b.numel()) ? b.data_ptr() : NULL;
|
||||
p.xref = (xref.numel()) ? xref.data_ptr() : NULL;
|
||||
p.yref = (yref.numel()) ? yref.data_ptr() : NULL;
|
||||
p.dy = (dy.numel()) ? dy.data_ptr() : NULL;
|
||||
p.y = y.data_ptr();
|
||||
p.grad = grad;
|
||||
p.act = act;
|
||||
p.alpha = alpha;
|
||||
p.gain = gain;
|
||||
p.clamp = clamp;
|
||||
p.sizeX = (int)x.numel();
|
||||
p.sizeB = (int)b.numel();
|
||||
p.stepB = (b.numel()) ? (int)x.stride(dim) : 1;
|
||||
|
||||
// Choose CUDA kernel.
|
||||
void* kernel;
|
||||
AT_DISPATCH_FLOATING_TYPES_AND_HALF(x.scalar_type(), "upfirdn2d_cuda", [&]
|
||||
{
|
||||
kernel = choose_bias_act_kernel<scalar_t>(p);
|
||||
});
|
||||
TORCH_CHECK(kernel, "no CUDA kernel found for the specified activation func");
|
||||
|
||||
// Launch CUDA kernel.
|
||||
p.loopX = 4;
|
||||
int blockSize = 4 * 32;
|
||||
int gridSize = (p.sizeX - 1) / (p.loopX * blockSize) + 1;
|
||||
void* args[] = {&p};
|
||||
AT_CUDA_CHECK(cudaLaunchKernel(kernel, gridSize, blockSize, args, 0, at::cuda::getCurrentCUDAStream()));
|
||||
return y;
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
|
||||
{
|
||||
m.def("bias_act", &bias_act);
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
@@ -0,0 +1,177 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
*
|
||||
* NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
* property and proprietary rights in and to this material, related
|
||||
* documentation and any modifications thereto. Any use, reproduction,
|
||||
* disclosure or distribution of this material and related documentation
|
||||
* without an express license agreement from NVIDIA CORPORATION or
|
||||
* its affiliates is strictly prohibited.
|
||||
*/
|
||||
|
||||
#include <c10/util/Half.h>
|
||||
#include "bias_act.h"
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// Helpers.
|
||||
|
||||
template <class T> struct InternalType;
|
||||
template <> struct InternalType<double> { typedef double scalar_t; };
|
||||
template <> struct InternalType<float> { typedef float scalar_t; };
|
||||
template <> struct InternalType<c10::Half> { typedef float scalar_t; };
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// CUDA kernel.
|
||||
|
||||
template <class T, int A>
|
||||
__global__ void bias_act_kernel(bias_act_kernel_params p)
|
||||
{
|
||||
typedef typename InternalType<T>::scalar_t scalar_t;
|
||||
int G = p.grad;
|
||||
scalar_t alpha = (scalar_t)p.alpha;
|
||||
scalar_t gain = (scalar_t)p.gain;
|
||||
scalar_t clamp = (scalar_t)p.clamp;
|
||||
scalar_t one = (scalar_t)1;
|
||||
scalar_t two = (scalar_t)2;
|
||||
scalar_t expRange = (scalar_t)80;
|
||||
scalar_t halfExpRange = (scalar_t)40;
|
||||
scalar_t seluScale = (scalar_t)1.0507009873554804934193349852946;
|
||||
scalar_t seluAlpha = (scalar_t)1.6732632423543772848170429916717;
|
||||
|
||||
// Loop over elements.
|
||||
int xi = blockIdx.x * p.loopX * blockDim.x + threadIdx.x;
|
||||
for (int loopIdx = 0; loopIdx < p.loopX && xi < p.sizeX; loopIdx++, xi += blockDim.x)
|
||||
{
|
||||
// Load.
|
||||
scalar_t x = (scalar_t)((const T*)p.x)[xi];
|
||||
scalar_t b = (p.b) ? (scalar_t)((const T*)p.b)[(xi / p.stepB) % p.sizeB] : 0;
|
||||
scalar_t xref = (p.xref) ? (scalar_t)((const T*)p.xref)[xi] : 0;
|
||||
scalar_t yref = (p.yref) ? (scalar_t)((const T*)p.yref)[xi] : 0;
|
||||
scalar_t dy = (p.dy) ? (scalar_t)((const T*)p.dy)[xi] : one;
|
||||
scalar_t yy = (gain != 0) ? yref / gain : 0;
|
||||
scalar_t y = 0;
|
||||
|
||||
// Apply bias.
|
||||
((G == 0) ? x : xref) += b;
|
||||
|
||||
// linear
|
||||
if (A == 1)
|
||||
{
|
||||
if (G == 0) y = x;
|
||||
if (G == 1) y = x;
|
||||
}
|
||||
|
||||
// relu
|
||||
if (A == 2)
|
||||
{
|
||||
if (G == 0) y = (x > 0) ? x : 0;
|
||||
if (G == 1) y = (yy > 0) ? x : 0;
|
||||
}
|
||||
|
||||
// lrelu
|
||||
if (A == 3)
|
||||
{
|
||||
if (G == 0) y = (x > 0) ? x : x * alpha;
|
||||
if (G == 1) y = (yy > 0) ? x : x * alpha;
|
||||
}
|
||||
|
||||
// tanh
|
||||
if (A == 4)
|
||||
{
|
||||
if (G == 0) { scalar_t c = exp(x); scalar_t d = one / c; y = (x < -expRange) ? -one : (x > expRange) ? one : (c - d) / (c + d); }
|
||||
if (G == 1) y = x * (one - yy * yy);
|
||||
if (G == 2) y = x * (one - yy * yy) * (-two * yy);
|
||||
}
|
||||
|
||||
// sigmoid
|
||||
if (A == 5)
|
||||
{
|
||||
if (G == 0) y = (x < -expRange) ? 0 : one / (exp(-x) + one);
|
||||
if (G == 1) y = x * yy * (one - yy);
|
||||
if (G == 2) y = x * yy * (one - yy) * (one - two * yy);
|
||||
}
|
||||
|
||||
// elu
|
||||
if (A == 6)
|
||||
{
|
||||
if (G == 0) y = (x >= 0) ? x : exp(x) - one;
|
||||
if (G == 1) y = (yy >= 0) ? x : x * (yy + one);
|
||||
if (G == 2) y = (yy >= 0) ? 0 : x * (yy + one);
|
||||
}
|
||||
|
||||
// selu
|
||||
if (A == 7)
|
||||
{
|
||||
if (G == 0) y = (x >= 0) ? seluScale * x : (seluScale * seluAlpha) * (exp(x) - one);
|
||||
if (G == 1) y = (yy >= 0) ? x * seluScale : x * (yy + seluScale * seluAlpha);
|
||||
if (G == 2) y = (yy >= 0) ? 0 : x * (yy + seluScale * seluAlpha);
|
||||
}
|
||||
|
||||
// softplus
|
||||
if (A == 8)
|
||||
{
|
||||
if (G == 0) y = (x > expRange) ? x : log(exp(x) + one);
|
||||
if (G == 1) y = x * (one - exp(-yy));
|
||||
if (G == 2) { scalar_t c = exp(-yy); y = x * c * (one - c); }
|
||||
}
|
||||
|
||||
// swish
|
||||
if (A == 9)
|
||||
{
|
||||
if (G == 0)
|
||||
y = (x < -expRange) ? 0 : x / (exp(-x) + one);
|
||||
else
|
||||
{
|
||||
scalar_t c = exp(xref);
|
||||
scalar_t d = c + one;
|
||||
if (G == 1)
|
||||
y = (xref > halfExpRange) ? x : x * c * (xref + d) / (d * d);
|
||||
else
|
||||
y = (xref > halfExpRange) ? 0 : x * c * (xref * (two - d) + two * d) / (d * d * d);
|
||||
yref = (xref < -expRange) ? 0 : xref / (exp(-xref) + one) * gain;
|
||||
}
|
||||
}
|
||||
|
||||
// Apply gain.
|
||||
y *= gain * dy;
|
||||
|
||||
// Clamp.
|
||||
if (clamp >= 0)
|
||||
{
|
||||
if (G == 0)
|
||||
y = (y > -clamp & y < clamp) ? y : (y >= 0) ? clamp : -clamp;
|
||||
else
|
||||
y = (yref > -clamp & yref < clamp) ? y : 0;
|
||||
}
|
||||
|
||||
// Store.
|
||||
((T*)p.y)[xi] = (T)y;
|
||||
}
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// CUDA kernel selection.
|
||||
|
||||
template <class T> void* choose_bias_act_kernel(const bias_act_kernel_params& p)
|
||||
{
|
||||
if (p.act == 1) return (void*)bias_act_kernel<T, 1>;
|
||||
if (p.act == 2) return (void*)bias_act_kernel<T, 2>;
|
||||
if (p.act == 3) return (void*)bias_act_kernel<T, 3>;
|
||||
if (p.act == 4) return (void*)bias_act_kernel<T, 4>;
|
||||
if (p.act == 5) return (void*)bias_act_kernel<T, 5>;
|
||||
if (p.act == 6) return (void*)bias_act_kernel<T, 6>;
|
||||
if (p.act == 7) return (void*)bias_act_kernel<T, 7>;
|
||||
if (p.act == 8) return (void*)bias_act_kernel<T, 8>;
|
||||
if (p.act == 9) return (void*)bias_act_kernel<T, 9>;
|
||||
return NULL;
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// Template specializations.
|
||||
|
||||
template void* choose_bias_act_kernel<double> (const bias_act_kernel_params& p);
|
||||
template void* choose_bias_act_kernel<float> (const bias_act_kernel_params& p);
|
||||
template void* choose_bias_act_kernel<c10::Half> (const bias_act_kernel_params& p);
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
@@ -0,0 +1,42 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
*
|
||||
* NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
* property and proprietary rights in and to this material, related
|
||||
* documentation and any modifications thereto. Any use, reproduction,
|
||||
* disclosure or distribution of this material and related documentation
|
||||
* without an express license agreement from NVIDIA CORPORATION or
|
||||
* its affiliates is strictly prohibited.
|
||||
*/
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// CUDA kernel parameters.
|
||||
|
||||
struct bias_act_kernel_params
|
||||
{
|
||||
const void* x; // [sizeX]
|
||||
const void* b; // [sizeB] or NULL
|
||||
const void* xref; // [sizeX] or NULL
|
||||
const void* yref; // [sizeX] or NULL
|
||||
const void* dy; // [sizeX] or NULL
|
||||
void* y; // [sizeX]
|
||||
|
||||
int grad;
|
||||
int act;
|
||||
float alpha;
|
||||
float gain;
|
||||
float clamp;
|
||||
|
||||
int sizeX;
|
||||
int sizeB;
|
||||
int stepB;
|
||||
int loopX;
|
||||
};
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// CUDA kernel selection.
|
||||
|
||||
template <class T> void* choose_bias_act_kernel(const bias_act_kernel_params& p);
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
@@ -0,0 +1,289 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""Custom PyTorch ops for efficient bias and activation."""
|
||||
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from ... import dnnlib
|
||||
from .. import custom_ops, misc
|
||||
|
||||
activation_funcs = {
|
||||
'linear':
|
||||
dnnlib.EasyDict(
|
||||
func=lambda x, **_: x,
|
||||
def_alpha=0,
|
||||
def_gain=1,
|
||||
cuda_idx=1,
|
||||
ref='',
|
||||
has_2nd_grad=False),
|
||||
'relu':
|
||||
dnnlib.EasyDict(
|
||||
func=lambda x, **_: torch.nn.functional.relu(x),
|
||||
def_alpha=0,
|
||||
def_gain=np.sqrt(2),
|
||||
cuda_idx=2,
|
||||
ref='y',
|
||||
has_2nd_grad=False),
|
||||
'lrelu':
|
||||
dnnlib.EasyDict(
|
||||
func=lambda x, alpha, **_: torch.nn.functional.leaky_relu(x, alpha),
|
||||
def_alpha=0.2,
|
||||
def_gain=np.sqrt(2),
|
||||
cuda_idx=3,
|
||||
ref='y',
|
||||
has_2nd_grad=False),
|
||||
'tanh':
|
||||
dnnlib.EasyDict(
|
||||
func=lambda x, **_: torch.tanh(x),
|
||||
def_alpha=0,
|
||||
def_gain=1,
|
||||
cuda_idx=4,
|
||||
ref='y',
|
||||
has_2nd_grad=True),
|
||||
'sigmoid':
|
||||
dnnlib.EasyDict(
|
||||
func=lambda x, **_: torch.sigmoid(x),
|
||||
def_alpha=0,
|
||||
def_gain=1,
|
||||
cuda_idx=5,
|
||||
ref='y',
|
||||
has_2nd_grad=True),
|
||||
'elu':
|
||||
dnnlib.EasyDict(
|
||||
func=lambda x, **_: torch.nn.functional.elu(x),
|
||||
def_alpha=0,
|
||||
def_gain=1,
|
||||
cuda_idx=6,
|
||||
ref='y',
|
||||
has_2nd_grad=True),
|
||||
'selu':
|
||||
dnnlib.EasyDict(
|
||||
func=lambda x, **_: torch.nn.functional.selu(x),
|
||||
def_alpha=0,
|
||||
def_gain=1,
|
||||
cuda_idx=7,
|
||||
ref='y',
|
||||
has_2nd_grad=True),
|
||||
'softplus':
|
||||
dnnlib.EasyDict(
|
||||
func=lambda x, **_: torch.nn.functional.softplus(x),
|
||||
def_alpha=0,
|
||||
def_gain=1,
|
||||
cuda_idx=8,
|
||||
ref='y',
|
||||
has_2nd_grad=True),
|
||||
'swish':
|
||||
dnnlib.EasyDict(
|
||||
func=lambda x, **_: torch.sigmoid(x) * x,
|
||||
def_alpha=0,
|
||||
def_gain=np.sqrt(2),
|
||||
cuda_idx=9,
|
||||
ref='x',
|
||||
has_2nd_grad=True),
|
||||
}
|
||||
|
||||
_plugin = None
|
||||
_null_tensor = torch.empty([0])
|
||||
|
||||
|
||||
def _init():
|
||||
global _plugin
|
||||
if _plugin is None:
|
||||
_plugin = custom_ops.get_plugin(
|
||||
module_name='bias_act_plugin',
|
||||
sources=['bias_act.cpp', 'bias_act.cu'],
|
||||
headers=['bias_act.h'],
|
||||
source_dir=os.path.dirname(__file__),
|
||||
extra_cuda_cflags=['--use_fast_math'],
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def bias_act(x,
|
||||
b=None,
|
||||
dim=1,
|
||||
act='linear',
|
||||
alpha=None,
|
||||
gain=None,
|
||||
clamp=None,
|
||||
impl='cuda'):
|
||||
r"""Fused bias and activation function.
|
||||
|
||||
Adds bias `b` to activation tensor `x`, evaluates activation function `act`,
|
||||
and scales the result by `gain`. Each of the steps is optional. In most cases,
|
||||
the fused op is considerably more efficient than performing the same calculation
|
||||
using standard PyTorch ops. It supports first and second order gradients,
|
||||
but not third order gradients.
|
||||
|
||||
Args:
|
||||
x: Input activation tensor. Can be of any shape.
|
||||
b: Bias vector, or `None` to disable. Must be a 1D tensor of the same type
|
||||
as `x`. The shape must be known, and it must match the dimension of `x`
|
||||
corresponding to `dim`.
|
||||
dim: The dimension in `x` corresponding to the elements of `b`.
|
||||
The value of `dim` is ignored if `b` is not specified.
|
||||
act: Name of the activation function to evaluate, or `"linear"` to disable.
|
||||
Can be e.g. `"relu"`, `"lrelu"`, `"tanh"`, `"sigmoid"`, `"swish"`, etc.
|
||||
See `activation_funcs` for a full list. `None` is not allowed.
|
||||
alpha: Shape parameter for the activation function, or `None` to use the default.
|
||||
gain: Scaling factor for the output tensor, or `None` to use default.
|
||||
See `activation_funcs` for the default scaling of each activation function.
|
||||
If unsure, consider specifying 1.
|
||||
clamp: Clamp the output values to `[-clamp, +clamp]`, or `None` to disable
|
||||
the clamping (default).
|
||||
impl: Name of the implementation to use. Can be `"ref"` or `"cuda"` (default).
|
||||
|
||||
Returns:
|
||||
Tensor of the same shape and datatype as `x`.
|
||||
"""
|
||||
assert isinstance(x, torch.Tensor)
|
||||
assert impl in ['ref', 'cuda']
|
||||
if impl == 'cuda' and x.device.type == 'cuda' and _init():
|
||||
return _bias_act_cuda(
|
||||
dim=dim, act=act, alpha=alpha, gain=gain, clamp=clamp).apply(x, b)
|
||||
return _bias_act_ref(
|
||||
x=x, b=b, dim=dim, act=act, alpha=alpha, gain=gain, clamp=clamp)
|
||||
|
||||
|
||||
@misc.profiled_function
|
||||
def _bias_act_ref(x,
|
||||
b=None,
|
||||
dim=1,
|
||||
act='linear',
|
||||
alpha=None,
|
||||
gain=None,
|
||||
clamp=None):
|
||||
"""Slow reference implementation of `bias_act()` using standard TensorFlow ops.
|
||||
"""
|
||||
assert isinstance(x, torch.Tensor)
|
||||
assert clamp is None or clamp >= 0
|
||||
spec = activation_funcs[act]
|
||||
alpha = float(alpha if alpha is not None else spec.def_alpha)
|
||||
gain = float(gain if gain is not None else spec.def_gain)
|
||||
clamp = float(clamp if clamp is not None else -1)
|
||||
|
||||
# Add bias.
|
||||
if b is not None:
|
||||
assert isinstance(b, torch.Tensor) and b.ndim == 1
|
||||
assert 0 <= dim < x.ndim
|
||||
assert b.shape[0] == x.shape[dim]
|
||||
x = x + b.reshape([-1 if i == dim else 1 for i in range(x.ndim)])
|
||||
|
||||
# Evaluate activation function.
|
||||
alpha = float(alpha)
|
||||
x = spec.func(x, alpha=alpha)
|
||||
|
||||
# Scale by gain.
|
||||
gain = float(gain)
|
||||
if gain != 1:
|
||||
x = x * gain
|
||||
|
||||
# Clamp.
|
||||
if clamp >= 0:
|
||||
x = x.clamp(-clamp, clamp) # pylint: disable=invalid-unary-operand-type
|
||||
return x
|
||||
|
||||
|
||||
_bias_act_cuda_cache = dict()
|
||||
|
||||
|
||||
def _bias_act_cuda(dim=1, act='linear', alpha=None, gain=None, clamp=None):
|
||||
"""Fast CUDA implementation of `bias_act()` using custom ops.
|
||||
"""
|
||||
# Parse arguments.
|
||||
assert clamp is None or clamp >= 0
|
||||
spec = activation_funcs[act]
|
||||
alpha = float(alpha if alpha is not None else spec.def_alpha)
|
||||
gain = float(gain if gain is not None else spec.def_gain)
|
||||
clamp = float(clamp if clamp is not None else -1)
|
||||
|
||||
# Lookup from cache.
|
||||
key = (dim, act, alpha, gain, clamp)
|
||||
if key in _bias_act_cuda_cache:
|
||||
return _bias_act_cuda_cache[key]
|
||||
|
||||
# Forward op.
|
||||
class BiasActCuda(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, x, b): # pylint: disable=arguments-differ
|
||||
ctx.memory_format = torch.channels_last if x.ndim > 2 and x.stride(
|
||||
1) == 1 else torch.contiguous_format
|
||||
x = x.contiguous(memory_format=ctx.memory_format)
|
||||
b = b.contiguous() if b is not None else _null_tensor
|
||||
y = x
|
||||
if act != 'linear' or gain != 1 or clamp >= 0 or b is not _null_tensor:
|
||||
y = _plugin.bias_act(x, b, _null_tensor, _null_tensor,
|
||||
_null_tensor, 0, dim, spec.cuda_idx,
|
||||
alpha, gain, clamp)
|
||||
ctx.save_for_backward(
|
||||
x if 'x' in spec.ref or spec.has_2nd_grad else _null_tensor,
|
||||
b if 'x' in spec.ref or spec.has_2nd_grad else _null_tensor,
|
||||
y if 'y' in spec.ref else _null_tensor)
|
||||
return y
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, dy): # pylint: disable=arguments-differ
|
||||
dy = dy.contiguous(memory_format=ctx.memory_format)
|
||||
x, b, y = ctx.saved_tensors
|
||||
dx = None
|
||||
db = None
|
||||
|
||||
if ctx.needs_input_grad[0] or ctx.needs_input_grad[1]:
|
||||
dx = dy
|
||||
if act != 'linear' or gain != 1 or clamp >= 0:
|
||||
dx = BiasActCudaGrad.apply(dy, x, b, y)
|
||||
|
||||
if ctx.needs_input_grad[1]:
|
||||
db = dx.sum([i for i in range(dx.ndim) if i != dim])
|
||||
|
||||
return dx, db
|
||||
|
||||
# Backward op.
|
||||
class BiasActCudaGrad(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, dy, x, b, y): # pylint: disable=arguments-differ
|
||||
ctx.memory_format = torch.channels_last if dy.ndim > 2 and dy.stride(
|
||||
1) == 1 else torch.contiguous_format
|
||||
dx = _plugin.bias_act(dy, b, x, y, _null_tensor, 1, dim,
|
||||
spec.cuda_idx, alpha, gain, clamp)
|
||||
ctx.save_for_backward(dy if spec.has_2nd_grad else _null_tensor, x,
|
||||
b, y)
|
||||
return dx
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, d_dx): # pylint: disable=arguments-differ
|
||||
d_dx = d_dx.contiguous(memory_format=ctx.memory_format)
|
||||
dy, x, b, y = ctx.saved_tensors
|
||||
d_dy = None
|
||||
d_x = None
|
||||
d_b = None
|
||||
d_y = None
|
||||
|
||||
if ctx.needs_input_grad[0]:
|
||||
d_dy = BiasActCudaGrad.apply(d_dx, x, b, y)
|
||||
|
||||
if spec.has_2nd_grad and (ctx.needs_input_grad[1]
|
||||
or ctx.needs_input_grad[2]):
|
||||
d_x = _plugin.bias_act(d_dx, b, x, y, dy, 2, dim,
|
||||
spec.cuda_idx, alpha, gain, clamp)
|
||||
|
||||
if spec.has_2nd_grad and ctx.needs_input_grad[2]:
|
||||
d_b = d_x.sum([i for i in range(d_x.ndim) if i != dim])
|
||||
|
||||
return d_dy, d_x, d_b, d_y
|
||||
|
||||
# Add to cache.
|
||||
_bias_act_cuda_cache[key] = BiasActCuda
|
||||
return BiasActCuda
|
||||
@@ -0,0 +1,296 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""Custom replacement for `torch.nn.functional.conv2d` that supports
|
||||
arbitrarily high order gradients with zero performance penalty."""
|
||||
|
||||
import contextlib
|
||||
|
||||
import torch
|
||||
|
||||
# pylint: disable=redefined-builtin
|
||||
# pylint: disable=arguments-differ
|
||||
# pylint: disable=protected-access
|
||||
|
||||
enabled = False # Enable the custom op by setting this to true.
|
||||
weight_gradients_disabled = False # Forcefully disable computation of gradients with respect to the weights.
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def no_weight_gradients(disable=True):
|
||||
global weight_gradients_disabled
|
||||
old = weight_gradients_disabled
|
||||
if disable:
|
||||
weight_gradients_disabled = True
|
||||
yield
|
||||
weight_gradients_disabled = old
|
||||
|
||||
|
||||
def conv2d(input,
|
||||
weight,
|
||||
bias=None,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1):
|
||||
if _should_use_custom_op(input):
|
||||
return _conv2d_gradfix(
|
||||
transpose=False,
|
||||
weight_shape=weight.shape,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
output_padding=0,
|
||||
dilation=dilation,
|
||||
groups=groups).apply(input, weight, bias)
|
||||
return torch.nn.functional.conv2d(
|
||||
input=input,
|
||||
weight=weight,
|
||||
bias=bias,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
groups=groups)
|
||||
|
||||
|
||||
def conv_transpose2d(input,
|
||||
weight,
|
||||
bias=None,
|
||||
stride=1,
|
||||
padding=0,
|
||||
output_padding=0,
|
||||
groups=1,
|
||||
dilation=1):
|
||||
if _should_use_custom_op(input):
|
||||
return _conv2d_gradfix(
|
||||
transpose=True,
|
||||
weight_shape=weight.shape,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
output_padding=output_padding,
|
||||
groups=groups,
|
||||
dilation=dilation).apply(input, weight, bias)
|
||||
return torch.nn.functional.conv_transpose2d(
|
||||
input=input,
|
||||
weight=weight,
|
||||
bias=bias,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
output_padding=output_padding,
|
||||
groups=groups,
|
||||
dilation=dilation)
|
||||
|
||||
|
||||
def _should_use_custom_op(input):
|
||||
assert isinstance(input, torch.Tensor)
|
||||
if (not enabled) or (not torch.backends.cudnn.enabled):
|
||||
return False
|
||||
if input.device.type != 'cuda':
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _tuple_of_ints(xs, ndim):
|
||||
xs = tuple(xs) if isinstance(xs, (tuple, list)) else (xs, ) * ndim
|
||||
assert len(xs) == ndim
|
||||
assert all(isinstance(x, int) for x in xs)
|
||||
return xs
|
||||
|
||||
|
||||
_conv2d_gradfix_cache = dict()
|
||||
_null_tensor = torch.empty([0])
|
||||
|
||||
|
||||
def _conv2d_gradfix(transpose, weight_shape, stride, padding, output_padding,
|
||||
dilation, groups):
|
||||
# Parse arguments.
|
||||
ndim = 2
|
||||
weight_shape = tuple(weight_shape)
|
||||
stride = _tuple_of_ints(stride, ndim)
|
||||
padding = _tuple_of_ints(padding, ndim)
|
||||
output_padding = _tuple_of_ints(output_padding, ndim)
|
||||
dilation = _tuple_of_ints(dilation, ndim)
|
||||
|
||||
# Lookup from cache.
|
||||
key = (transpose, weight_shape, stride, padding, output_padding, dilation,
|
||||
groups)
|
||||
if key in _conv2d_gradfix_cache:
|
||||
return _conv2d_gradfix_cache[key]
|
||||
|
||||
# Validate arguments.
|
||||
assert groups >= 1
|
||||
assert len(weight_shape) == ndim + 2
|
||||
assert all(stride[i] >= 1 for i in range(ndim))
|
||||
assert all(padding[i] >= 0 for i in range(ndim))
|
||||
assert all(dilation[i] >= 0 for i in range(ndim))
|
||||
if not transpose:
|
||||
assert all(output_padding[i] == 0 for i in range(ndim))
|
||||
else: # transpose
|
||||
assert all(0 <= output_padding[i] < max(stride[i], dilation[i])
|
||||
for i in range(ndim))
|
||||
|
||||
# Helpers.
|
||||
common_kwargs = dict(
|
||||
stride=stride, padding=padding, dilation=dilation, groups=groups)
|
||||
|
||||
def calc_output_padding(input_shape, output_shape):
|
||||
if transpose:
|
||||
return [0, 0]
|
||||
|
||||
result_list = []
|
||||
for i in range(ndim):
|
||||
temp1 = input_shape[i + 2]
|
||||
temp2 = (output_shape[i + 2] - 1) * stride[i]
|
||||
temp3 = (1 - 2 * padding[i])
|
||||
temp4 = dilation[i] * (weight_shape[i + 2] - 1)
|
||||
result = temp1 - temp2 - temp3 - temp4
|
||||
result_list.append(result)
|
||||
|
||||
return result_list
|
||||
|
||||
# Forward & backward.
|
||||
class Conv2d(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input, weight, bias):
|
||||
assert weight.shape == weight_shape
|
||||
ctx.save_for_backward(
|
||||
input if weight.requires_grad else _null_tensor,
|
||||
weight if input.requires_grad else _null_tensor,
|
||||
)
|
||||
ctx.input_shape = input.shape
|
||||
|
||||
# Simple 1x1 convolution => cuBLAS (only on Volta, not on Ampere).
|
||||
if weight_shape[2:] == stride == dilation == (
|
||||
1, 1) and padding == (
|
||||
0, 0) and torch.cuda.get_device_capability(
|
||||
input.device) < (8, 0):
|
||||
a = weight.reshape(groups, weight_shape[0] // groups,
|
||||
weight_shape[1])
|
||||
b = input.reshape(input.shape[0], groups,
|
||||
input.shape[1] // groups, -1)
|
||||
c = (a.transpose(1, 2) if transpose else a) @ b.permute(
|
||||
1, 2, 0, 3).flatten(2)
|
||||
c = c.reshape(-1, input.shape[0],
|
||||
*input.shape[2:]).transpose(0, 1)
|
||||
c = c if bias is None else c + bias.unsqueeze(0).unsqueeze(
|
||||
2).unsqueeze(3)
|
||||
if input.stride(1) == 1:
|
||||
return c.contiguous(memory_format=torch.channels_last)
|
||||
else:
|
||||
return c.contiguous(memory_format=torch.contiguous_format)
|
||||
# General case => cuDNN.
|
||||
if transpose:
|
||||
return torch.nn.functional.conv_transpose2d(
|
||||
input=input,
|
||||
weight=weight,
|
||||
bias=bias,
|
||||
output_padding=output_padding,
|
||||
**common_kwargs)
|
||||
return torch.nn.functional.conv2d(
|
||||
input=input, weight=weight, bias=bias, **common_kwargs)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
input, weight = ctx.saved_tensors
|
||||
input_shape = ctx.input_shape
|
||||
grad_input = None
|
||||
grad_weight = None
|
||||
grad_bias = None
|
||||
|
||||
if ctx.needs_input_grad[0]:
|
||||
p = calc_output_padding(
|
||||
input_shape=input_shape, output_shape=grad_output.shape)
|
||||
op = _conv2d_gradfix(
|
||||
transpose=(not transpose),
|
||||
weight_shape=weight_shape,
|
||||
output_padding=p,
|
||||
**common_kwargs)
|
||||
grad_input = op.apply(grad_output, weight, None)
|
||||
assert grad_input.shape == input_shape
|
||||
|
||||
if ctx.needs_input_grad[1] and not weight_gradients_disabled:
|
||||
grad_weight = Conv2dGradWeight.apply(grad_output, input,
|
||||
weight)
|
||||
assert grad_weight.shape == weight_shape
|
||||
|
||||
if ctx.needs_input_grad[2]:
|
||||
grad_bias = grad_output.sum([0, 2, 3])
|
||||
|
||||
return grad_input, grad_weight, grad_bias
|
||||
|
||||
# Gradient with respect to the weights.
|
||||
class Conv2dGradWeight(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, grad_output, input, weight):
|
||||
ctx.save_for_backward(
|
||||
grad_output if input.requires_grad else _null_tensor,
|
||||
input if grad_output.requires_grad else _null_tensor,
|
||||
)
|
||||
ctx.grad_output_shape = grad_output.shape
|
||||
ctx.input_shape = input.shape
|
||||
|
||||
# Simple 1x1 convolution => cuBLAS (on both Volta and Ampere).
|
||||
if weight_shape[2:] == stride == dilation == (
|
||||
1, 1) and padding == (0, 0):
|
||||
a = grad_output.reshape(grad_output.shape[0], groups,
|
||||
grad_output.shape[1] // groups,
|
||||
-1).permute(1, 2, 0, 3).flatten(2)
|
||||
b = input.reshape(input.shape[0], groups,
|
||||
input.shape[1] // groups,
|
||||
-1).permute(1, 2, 0, 3).flatten(2)
|
||||
c = (b @ a.transpose(1, 2) if transpose else a
|
||||
@ b.transpose(1, 2)).reshape(weight_shape)
|
||||
if input.stride(1) == 1:
|
||||
return c.contiguous(memory_format=torch.channels_last)
|
||||
else:
|
||||
return c.contiguous(memory_format=torch.contiguous_format)
|
||||
|
||||
# General case => cuDNN.
|
||||
return torch.ops.aten.convolution_backward(
|
||||
grad_output=grad_output,
|
||||
input=input,
|
||||
weight=weight,
|
||||
bias_sizes=None,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
transposed=transpose,
|
||||
output_padding=output_padding,
|
||||
groups=groups,
|
||||
output_mask=[False, True, False])[1]
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad2_grad_weight):
|
||||
grad_output, input = ctx.saved_tensors
|
||||
grad_output_shape = ctx.grad_output_shape
|
||||
input_shape = ctx.input_shape
|
||||
grad2_grad_output = None
|
||||
grad2_input = None
|
||||
|
||||
if ctx.needs_input_grad[0]:
|
||||
grad2_grad_output = Conv2d.apply(input, grad2_grad_weight,
|
||||
None)
|
||||
assert grad2_grad_output.shape == grad_output_shape
|
||||
|
||||
if ctx.needs_input_grad[1]:
|
||||
p = calc_output_padding(
|
||||
input_shape=input_shape, output_shape=grad_output_shape)
|
||||
op = _conv2d_gradfix(
|
||||
transpose=(not transpose),
|
||||
weight_shape=weight_shape,
|
||||
output_padding=p,
|
||||
**common_kwargs)
|
||||
grad2_input = op.apply(grad_output, grad2_grad_weight, None)
|
||||
assert grad2_input.shape == input_shape
|
||||
|
||||
return grad2_grad_output, grad2_input
|
||||
|
||||
_conv2d_gradfix_cache[key] = Conv2d
|
||||
return Conv2d
|
||||
@@ -0,0 +1,192 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""2D convolution with optional up/downsampling."""
|
||||
|
||||
import torch
|
||||
|
||||
from .. import misc
|
||||
from . import conv2d_gradfix, upfirdn2d
|
||||
from .upfirdn2d import _get_filter_size, _parse_padding
|
||||
|
||||
|
||||
def _get_weight_shape(w):
|
||||
with misc.suppress_tracer_warnings():
|
||||
shape = [int(sz) for sz in w.shape]
|
||||
misc.assert_shape(w, shape)
|
||||
return shape
|
||||
|
||||
|
||||
def _conv2d_wrapper(x,
|
||||
w,
|
||||
stride=1,
|
||||
padding=0,
|
||||
groups=1,
|
||||
transpose=False,
|
||||
flip_weight=True):
|
||||
"""Wrapper for the underlying `conv2d()` and `conv_transpose2d()` implementations.
|
||||
"""
|
||||
_out_channels, _in_channels_per_group, kh, kw = _get_weight_shape(w)
|
||||
|
||||
# Flip weight if requested.
|
||||
# Note: conv2d() actually performs correlation (flip_weight=True) not convolution (flip_weight=False).
|
||||
if not flip_weight and (kw > 1 or kh > 1):
|
||||
w = w.flip([2, 3])
|
||||
|
||||
# Execute using conv2d_gradfix.
|
||||
op = conv2d_gradfix.conv_transpose2d if transpose else conv2d_gradfix.conv2d
|
||||
return op(x, w, stride=stride, padding=padding, groups=groups)
|
||||
|
||||
|
||||
@misc.profiled_function
|
||||
def conv2d_resample(x,
|
||||
w,
|
||||
f=None,
|
||||
up=1,
|
||||
down=1,
|
||||
padding=0,
|
||||
groups=1,
|
||||
flip_weight=True,
|
||||
flip_filter=False):
|
||||
r"""2D convolution with optional up/downsampling.
|
||||
|
||||
Padding is performed only once at the beginning, not between the operations.
|
||||
|
||||
Args:
|
||||
x: Input tensor of shape
|
||||
`[batch_size, in_channels, in_height, in_width]`.
|
||||
w: Weight tensor of shape
|
||||
`[out_channels, in_channels//groups, kernel_height, kernel_width]`.
|
||||
f: Low-pass filter for up/downsampling. Must be prepared beforehand by
|
||||
calling upfirdn2d.setup_filter(). None = identity (default).
|
||||
up: Integer upsampling factor (default: 1).
|
||||
down: Integer downsampling factor (default: 1).
|
||||
padding: Padding with respect to the upsampled image. Can be a single number
|
||||
or a list/tuple `[x, y]` or `[x_before, x_after, y_before, y_after]`
|
||||
(default: 0).
|
||||
groups: Split input channels into N groups (default: 1).
|
||||
flip_weight: False = convolution, True = correlation (default: True).
|
||||
flip_filter: False = convolution, True = correlation (default: False).
|
||||
|
||||
Returns:
|
||||
Tensor of the shape `[batch_size, num_channels, out_height, out_width]`.
|
||||
"""
|
||||
# Validate arguments.
|
||||
assert isinstance(x, torch.Tensor) and (x.ndim == 4)
|
||||
assert isinstance(w, torch.Tensor) and (w.ndim == 4) and (w.dtype
|
||||
== x.dtype)
|
||||
assert f is None or (isinstance(f, torch.Tensor) and f.ndim in [1, 2]
|
||||
and f.dtype == torch.float32)
|
||||
assert isinstance(up, int) and (up >= 1)
|
||||
assert isinstance(down, int) and (down >= 1)
|
||||
assert isinstance(groups, int) and (groups >= 1)
|
||||
out_channels, in_channels_per_group, kh, kw = _get_weight_shape(w)
|
||||
fw, fh = _get_filter_size(f)
|
||||
px0, px1, py0, py1 = _parse_padding(padding)
|
||||
|
||||
# Adjust padding to account for up/downsampling.
|
||||
if up > 1:
|
||||
px0 += (fw + up - 1) // 2
|
||||
px1 += (fw - up) // 2
|
||||
py0 += (fh + up - 1) // 2
|
||||
py1 += (fh - up) // 2
|
||||
if down > 1:
|
||||
px0 += (fw - down + 1) // 2
|
||||
px1 += (fw - down) // 2
|
||||
py0 += (fh - down + 1) // 2
|
||||
py1 += (fh - down) // 2
|
||||
|
||||
# Fast path: 1x1 convolution with downsampling only => downsample first, then convolve.
|
||||
if kw == 1 and kh == 1 and (down > 1 and up == 1):
|
||||
x = upfirdn2d.upfirdn2d(
|
||||
x=x,
|
||||
f=f,
|
||||
down=down,
|
||||
padding=[px0, px1, py0, py1],
|
||||
flip_filter=flip_filter)
|
||||
x = _conv2d_wrapper(x=x, w=w, groups=groups, flip_weight=flip_weight)
|
||||
return x
|
||||
|
||||
# Fast path: 1x1 convolution with upsampling only => convolve first, then upsample.
|
||||
if kw == 1 and kh == 1 and (up > 1 and down == 1):
|
||||
x = _conv2d_wrapper(x=x, w=w, groups=groups, flip_weight=flip_weight)
|
||||
x = upfirdn2d.upfirdn2d(
|
||||
x=x,
|
||||
f=f,
|
||||
up=up,
|
||||
padding=[px0, px1, py0, py1],
|
||||
gain=up**2,
|
||||
flip_filter=flip_filter)
|
||||
return x
|
||||
|
||||
# Fast path: downsampling only => use strided convolution.
|
||||
if down > 1 and up == 1:
|
||||
x = upfirdn2d.upfirdn2d(
|
||||
x=x, f=f, padding=[px0, px1, py0, py1], flip_filter=flip_filter)
|
||||
x = _conv2d_wrapper(
|
||||
x=x, w=w, stride=down, groups=groups, flip_weight=flip_weight)
|
||||
return x
|
||||
|
||||
# Fast path: upsampling with optional downsampling => use transpose strided convolution.
|
||||
if up > 1:
|
||||
if groups == 1:
|
||||
w = w.transpose(0, 1)
|
||||
else:
|
||||
w = w.reshape(groups, out_channels // groups,
|
||||
in_channels_per_group, kh, kw)
|
||||
w = w.transpose(1, 2)
|
||||
w = w.reshape(groups * in_channels_per_group,
|
||||
out_channels // groups, kh, kw)
|
||||
px0 -= kw - 1
|
||||
px1 -= kw - up
|
||||
py0 -= kh - 1
|
||||
py1 -= kh - up
|
||||
pxt = max(min(-px0, -px1), 0)
|
||||
pyt = max(min(-py0, -py1), 0)
|
||||
x = _conv2d_wrapper(
|
||||
x=x,
|
||||
w=w,
|
||||
stride=up,
|
||||
padding=[pyt, pxt],
|
||||
groups=groups,
|
||||
transpose=True,
|
||||
flip_weight=(not flip_weight))
|
||||
x = upfirdn2d.upfirdn2d(
|
||||
x=x,
|
||||
f=f,
|
||||
padding=[px0 + pxt, px1 + pxt, py0 + pyt, py1 + pyt],
|
||||
gain=up**2,
|
||||
flip_filter=flip_filter)
|
||||
if down > 1:
|
||||
x = upfirdn2d.upfirdn2d(
|
||||
x=x, f=f, down=down, flip_filter=flip_filter)
|
||||
return x
|
||||
|
||||
# Fast path: no up/downsampling, padding supported by the underlying implementation => use plain conv2d.
|
||||
if up == 1 and down == 1:
|
||||
if px0 == px1 and py0 == py1 and px0 >= 0 and py0 >= 0:
|
||||
return _conv2d_wrapper(
|
||||
x=x,
|
||||
w=w,
|
||||
padding=[py0, px0],
|
||||
groups=groups,
|
||||
flip_weight=flip_weight)
|
||||
|
||||
# Fallback: Generic reference implementation.
|
||||
x = upfirdn2d.upfirdn2d(
|
||||
x=x,
|
||||
f=(f if up > 1 else None),
|
||||
up=up,
|
||||
padding=[px0, px1, py0, py1],
|
||||
gain=up**2,
|
||||
flip_filter=flip_filter)
|
||||
x = _conv2d_wrapper(x=x, w=w, groups=groups, flip_weight=flip_weight)
|
||||
if down > 1:
|
||||
x = upfirdn2d.upfirdn2d(x=x, f=f, down=down, flip_filter=flip_filter)
|
||||
return x
|
||||
@@ -0,0 +1,304 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
*
|
||||
* NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
* property and proprietary rights in and to this material, related
|
||||
* documentation and any modifications thereto. Any use, reproduction,
|
||||
* disclosure or distribution of this material and related documentation
|
||||
* without an express license agreement from NVIDIA CORPORATION or
|
||||
* its affiliates is strictly prohibited.
|
||||
*/
|
||||
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include "filtered_lrelu.h"
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
|
||||
static std::tuple<torch::Tensor, torch::Tensor, int> filtered_lrelu(
|
||||
torch::Tensor x, torch::Tensor fu, torch::Tensor fd, torch::Tensor b, torch::Tensor si,
|
||||
int up, int down, int px0, int px1, int py0, int py1, int sx, int sy, float gain, float slope, float clamp, bool flip_filters, bool writeSigns)
|
||||
{
|
||||
// Set CUDA device.
|
||||
TORCH_CHECK(x.is_cuda(), "x must reside on CUDA device");
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(x));
|
||||
|
||||
// Validate arguments.
|
||||
TORCH_CHECK(fu.device() == x.device() && fd.device() == x.device() && b.device() == x.device(), "all input tensors must reside on the same device");
|
||||
TORCH_CHECK(fu.dtype() == torch::kFloat && fd.dtype() == torch::kFloat, "fu and fd must be float32");
|
||||
TORCH_CHECK(b.dtype() == x.dtype(), "x and b must have the same dtype");
|
||||
TORCH_CHECK(x.dtype() == torch::kHalf || x.dtype() == torch::kFloat, "x and b must be float16 or float32");
|
||||
TORCH_CHECK(x.dim() == 4, "x must be rank 4");
|
||||
TORCH_CHECK(x.size(0) * x.size(1) <= INT_MAX && x.size(2) <= INT_MAX && x.size(3) <= INT_MAX, "x is too large");
|
||||
TORCH_CHECK(x.numel() > 0, "x is empty");
|
||||
TORCH_CHECK((fu.dim() == 1 || fu.dim() == 2) && (fd.dim() == 1 || fd.dim() == 2), "fu and fd must be rank 1 or 2");
|
||||
TORCH_CHECK(fu.size(0) <= INT_MAX && fu.size(-1) <= INT_MAX, "fu is too large");
|
||||
TORCH_CHECK(fd.size(0) <= INT_MAX && fd.size(-1) <= INT_MAX, "fd is too large");
|
||||
TORCH_CHECK(fu.numel() > 0, "fu is empty");
|
||||
TORCH_CHECK(fd.numel() > 0, "fd is empty");
|
||||
TORCH_CHECK(b.dim() == 1 && b.size(0) == x.size(1), "b must be a vector with the same number of channels as x");
|
||||
TORCH_CHECK(up >= 1 && down >= 1, "up and down must be at least 1");
|
||||
|
||||
// Figure out how much shared memory is available on the device.
|
||||
int maxSharedBytes = 0;
|
||||
AT_CUDA_CHECK(cudaDeviceGetAttribute(&maxSharedBytes, cudaDevAttrMaxSharedMemoryPerBlockOptin, x.device().index()));
|
||||
int sharedKB = maxSharedBytes >> 10;
|
||||
|
||||
// Populate enough launch parameters to check if a CUDA kernel exists.
|
||||
filtered_lrelu_kernel_params p;
|
||||
p.up = up;
|
||||
p.down = down;
|
||||
p.fuShape = make_int2((int)fu.size(-1), fu.dim() == 2 ? (int)fu.size(0) : 0); // shape [n, 0] indicates separable filter.
|
||||
p.fdShape = make_int2((int)fd.size(-1), fd.dim() == 2 ? (int)fd.size(0) : 0);
|
||||
filtered_lrelu_kernel_spec test_spec = choose_filtered_lrelu_kernel<float, int32_t, false, false>(p, sharedKB);
|
||||
if (!test_spec.exec)
|
||||
{
|
||||
// No kernel found - return empty tensors and indicate missing kernel with return code of -1.
|
||||
return std::make_tuple(torch::Tensor(), torch::Tensor(), -1);
|
||||
}
|
||||
|
||||
// Input/output element size.
|
||||
int64_t sz = (x.dtype() == torch::kHalf) ? 2 : 4;
|
||||
|
||||
// Input sizes.
|
||||
int64_t xw = (int)x.size(3);
|
||||
int64_t xh = (int)x.size(2);
|
||||
int64_t fut_w = (int)fu.size(-1) - 1;
|
||||
int64_t fut_h = (int)fu.size(0) - 1;
|
||||
int64_t fdt_w = (int)fd.size(-1) - 1;
|
||||
int64_t fdt_h = (int)fd.size(0) - 1;
|
||||
|
||||
// Logical size of upsampled buffer.
|
||||
int64_t cw = xw * up + (px0 + px1) - fut_w;
|
||||
int64_t ch = xh * up + (py0 + py1) - fut_h;
|
||||
TORCH_CHECK(cw > fdt_w && ch > fdt_h, "upsampled buffer must be at least the size of downsampling filter");
|
||||
TORCH_CHECK(cw <= INT_MAX && ch <= INT_MAX, "upsampled buffer is too large");
|
||||
|
||||
// Compute output size and allocate.
|
||||
int64_t yw = (cw - fdt_w + (down - 1)) / down;
|
||||
int64_t yh = (ch - fdt_h + (down - 1)) / down;
|
||||
TORCH_CHECK(yw > 0 && yh > 0, "output must be at least 1x1");
|
||||
TORCH_CHECK(yw <= INT_MAX && yh <= INT_MAX, "output is too large");
|
||||
torch::Tensor y = torch::empty({x.size(0), x.size(1), yh, yw}, x.options(), x.suggest_memory_format());
|
||||
|
||||
// Allocate sign tensor.
|
||||
torch::Tensor so;
|
||||
torch::Tensor s = si;
|
||||
bool readSigns = !!s.numel();
|
||||
int64_t sw_active = 0; // Active width of sign tensor.
|
||||
if (writeSigns)
|
||||
{
|
||||
sw_active = yw * down - (down - 1) + fdt_w; // Active width in elements.
|
||||
int64_t sh = yh * down - (down - 1) + fdt_h; // Height = active height.
|
||||
int64_t sw = (sw_active + 15) & ~15; // Width = active width in elements, rounded up to multiple of 16.
|
||||
TORCH_CHECK(sh <= INT_MAX && (sw >> 2) <= INT_MAX, "signs is too large");
|
||||
s = so = torch::empty({x.size(0), x.size(1), sh, sw >> 2}, x.options().dtype(torch::kUInt8), at::MemoryFormat::Contiguous);
|
||||
}
|
||||
else if (readSigns)
|
||||
sw_active = s.size(3) << 2;
|
||||
|
||||
// Validate sign tensor if in use.
|
||||
if (readSigns || writeSigns)
|
||||
{
|
||||
TORCH_CHECK(s.is_contiguous(), "signs must be contiguous");
|
||||
TORCH_CHECK(s.dtype() == torch::kUInt8, "signs must be uint8");
|
||||
TORCH_CHECK(s.device() == x.device(), "signs must reside on the same device as x");
|
||||
TORCH_CHECK(s.dim() == 4, "signs must be rank 4");
|
||||
TORCH_CHECK(s.size(0) == x.size(0) && s.size(1) == x.size(1), "signs must have same batch & channels as x");
|
||||
TORCH_CHECK(s.size(2) <= INT_MAX && s.size(3) <= INT_MAX, "signs is too large");
|
||||
}
|
||||
|
||||
// Populate rest of CUDA kernel parameters.
|
||||
p.x = x.data_ptr();
|
||||
p.y = y.data_ptr();
|
||||
p.b = b.data_ptr();
|
||||
p.s = (readSigns || writeSigns) ? s.data_ptr<unsigned char>() : 0;
|
||||
p.fu = fu.data_ptr<float>();
|
||||
p.fd = fd.data_ptr<float>();
|
||||
p.pad0 = make_int2(px0, py0);
|
||||
p.gain = gain;
|
||||
p.slope = slope;
|
||||
p.clamp = clamp;
|
||||
p.flip = (flip_filters) ? 1 : 0;
|
||||
p.xShape = make_int4((int)x.size(3), (int)x.size(2), (int)x.size(1), (int)x.size(0));
|
||||
p.yShape = make_int4((int)y.size(3), (int)y.size(2), (int)y.size(1), (int)y.size(0));
|
||||
p.sShape = (readSigns || writeSigns) ? make_int2((int)s.size(3), (int)s.size(2)) : make_int2(0, 0); // Width is in bytes. Contiguous.
|
||||
p.sOfs = make_int2(sx, sy);
|
||||
p.swLimit = (sw_active + 3) >> 2; // Rounded up to bytes.
|
||||
|
||||
// x, y, b strides are in bytes.
|
||||
p.xStride = make_longlong4(sz * x.stride(3), sz * x.stride(2), sz * x.stride(1), sz * x.stride(0));
|
||||
p.yStride = make_longlong4(sz * y.stride(3), sz * y.stride(2), sz * y.stride(1), sz * y.stride(0));
|
||||
p.bStride = sz * b.stride(0);
|
||||
|
||||
// fu, fd strides are in elements.
|
||||
p.fuStride = make_longlong3(fu.stride(-1), fu.dim() == 2 ? fu.stride(0) : 0, 0);
|
||||
p.fdStride = make_longlong3(fd.stride(-1), fd.dim() == 2 ? fd.stride(0) : 0, 0);
|
||||
|
||||
// Determine if indices don't fit in int32. Support negative strides although Torch currently never produces those.
|
||||
bool index64b = false;
|
||||
if (std::abs(p.bStride * x.size(1)) > INT_MAX) index64b = true;
|
||||
if (std::min(x.size(0) * p.xStride.w, 0ll) + std::min(x.size(1) * p.xStride.z, 0ll) + std::min(x.size(2) * p.xStride.y, 0ll) + std::min(x.size(3) * p.xStride.x, 0ll) < -INT_MAX) index64b = true;
|
||||
if (std::max(x.size(0) * p.xStride.w, 0ll) + std::max(x.size(1) * p.xStride.z, 0ll) + std::max(x.size(2) * p.xStride.y, 0ll) + std::max(x.size(3) * p.xStride.x, 0ll) > INT_MAX) index64b = true;
|
||||
if (std::min(y.size(0) * p.yStride.w, 0ll) + std::min(y.size(1) * p.yStride.z, 0ll) + std::min(y.size(2) * p.yStride.y, 0ll) + std::min(y.size(3) * p.yStride.x, 0ll) < -INT_MAX) index64b = true;
|
||||
if (std::max(y.size(0) * p.yStride.w, 0ll) + std::max(y.size(1) * p.yStride.z, 0ll) + std::max(y.size(2) * p.yStride.y, 0ll) + std::max(y.size(3) * p.yStride.x, 0ll) > INT_MAX) index64b = true;
|
||||
if (s.numel() > INT_MAX) index64b = true;
|
||||
|
||||
// Choose CUDA kernel.
|
||||
filtered_lrelu_kernel_spec spec = { 0 };
|
||||
AT_DISPATCH_FLOATING_TYPES_AND_HALF(x.scalar_type(), "filtered_lrelu_cuda", [&]
|
||||
{
|
||||
if constexpr (sizeof(scalar_t) <= 4) // Exclude doubles. constexpr prevents template instantiation.
|
||||
{
|
||||
// Choose kernel based on index type, datatype and sign read/write modes.
|
||||
if (!index64b && writeSigns && !readSigns) spec = choose_filtered_lrelu_kernel<scalar_t, int32_t, true, false>(p, sharedKB);
|
||||
else if (!index64b && !writeSigns && readSigns) spec = choose_filtered_lrelu_kernel<scalar_t, int32_t, false, true >(p, sharedKB);
|
||||
else if (!index64b && !writeSigns && !readSigns) spec = choose_filtered_lrelu_kernel<scalar_t, int32_t, false, false>(p, sharedKB);
|
||||
else if ( index64b && writeSigns && !readSigns) spec = choose_filtered_lrelu_kernel<scalar_t, int64_t, true, false>(p, sharedKB);
|
||||
else if ( index64b && !writeSigns && readSigns) spec = choose_filtered_lrelu_kernel<scalar_t, int64_t, false, true >(p, sharedKB);
|
||||
else if ( index64b && !writeSigns && !readSigns) spec = choose_filtered_lrelu_kernel<scalar_t, int64_t, false, false>(p, sharedKB);
|
||||
}
|
||||
});
|
||||
TORCH_CHECK(spec.exec, "internal error - CUDA kernel not found") // This should not happen because we tested earlier that kernel exists.
|
||||
|
||||
// Launch CUDA kernel.
|
||||
void* args[] = {&p};
|
||||
int bx = spec.numWarps * 32;
|
||||
int gx = (p.yShape.x - 1) / spec.tileOut.x + 1;
|
||||
int gy = (p.yShape.y - 1) / spec.tileOut.y + 1;
|
||||
int gz = p.yShape.z * p.yShape.w;
|
||||
|
||||
// Repeat multiple horizontal tiles in a CTA?
|
||||
if (spec.xrep)
|
||||
{
|
||||
p.tilesXrep = spec.xrep;
|
||||
p.tilesXdim = gx;
|
||||
|
||||
gx = (gx + p.tilesXrep - 1) / p.tilesXrep;
|
||||
std::swap(gx, gy);
|
||||
}
|
||||
else
|
||||
{
|
||||
p.tilesXrep = 0;
|
||||
p.tilesXdim = 0;
|
||||
}
|
||||
|
||||
// Launch filter setup kernel.
|
||||
AT_CUDA_CHECK(cudaLaunchKernel(spec.setup, 1, 1024, args, 0, at::cuda::getCurrentCUDAStream()));
|
||||
|
||||
// Copy kernels to constant memory.
|
||||
if ( writeSigns && !readSigns) AT_CUDA_CHECK((copy_filters<true, false>(at::cuda::getCurrentCUDAStream())));
|
||||
else if (!writeSigns && readSigns) AT_CUDA_CHECK((copy_filters<false, true >(at::cuda::getCurrentCUDAStream())));
|
||||
else if (!writeSigns && !readSigns) AT_CUDA_CHECK((copy_filters<false, false>(at::cuda::getCurrentCUDAStream())));
|
||||
|
||||
// Set cache and shared memory configurations for main kernel.
|
||||
AT_CUDA_CHECK(cudaFuncSetCacheConfig(spec.exec, cudaFuncCachePreferShared));
|
||||
if (spec.dynamicSharedKB) // Need dynamically allocated shared memory?
|
||||
AT_CUDA_CHECK(cudaFuncSetAttribute(spec.exec, cudaFuncAttributeMaxDynamicSharedMemorySize, spec.dynamicSharedKB << 10));
|
||||
AT_CUDA_CHECK(cudaFuncSetSharedMemConfig(spec.exec, cudaSharedMemBankSizeFourByte));
|
||||
|
||||
// Launch main kernel.
|
||||
const int maxSubGz = 65535; // CUDA maximum for block z dimension.
|
||||
for (int zofs=0; zofs < gz; zofs += maxSubGz) // Do multiple launches if gz is too big.
|
||||
{
|
||||
p.blockZofs = zofs;
|
||||
int subGz = std::min(maxSubGz, gz - zofs);
|
||||
AT_CUDA_CHECK(cudaLaunchKernel(spec.exec, dim3(gx, gy, subGz), bx, args, spec.dynamicSharedKB << 10, at::cuda::getCurrentCUDAStream()));
|
||||
}
|
||||
|
||||
// Done.
|
||||
return std::make_tuple(y, so, 0);
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
|
||||
static torch::Tensor filtered_lrelu_act(torch::Tensor x, torch::Tensor si, int sx, int sy, float gain, float slope, float clamp, bool writeSigns)
|
||||
{
|
||||
// Set CUDA device.
|
||||
TORCH_CHECK(x.is_cuda(), "x must reside on CUDA device");
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(x));
|
||||
|
||||
// Validate arguments.
|
||||
TORCH_CHECK(x.dim() == 4, "x must be rank 4");
|
||||
TORCH_CHECK(x.size(0) * x.size(1) <= INT_MAX && x.size(2) <= INT_MAX && x.size(3) <= INT_MAX, "x is too large");
|
||||
TORCH_CHECK(x.numel() > 0, "x is empty");
|
||||
TORCH_CHECK(x.dtype() == torch::kHalf || x.dtype() == torch::kFloat || x.dtype() == torch::kDouble, "x must be float16, float32 or float64");
|
||||
|
||||
// Output signs if we don't have sign input.
|
||||
torch::Tensor so;
|
||||
torch::Tensor s = si;
|
||||
bool readSigns = !!s.numel();
|
||||
if (writeSigns)
|
||||
{
|
||||
int64_t sw = x.size(3);
|
||||
sw = (sw + 15) & ~15; // Round to a multiple of 16 for coalescing.
|
||||
s = so = torch::empty({x.size(0), x.size(1), x.size(2), sw >> 2}, x.options().dtype(torch::kUInt8), at::MemoryFormat::Contiguous);
|
||||
}
|
||||
|
||||
// Validate sign tensor if in use.
|
||||
if (readSigns || writeSigns)
|
||||
{
|
||||
TORCH_CHECK(s.is_contiguous(), "signs must be contiguous");
|
||||
TORCH_CHECK(s.dtype() == torch::kUInt8, "signs must be uint8");
|
||||
TORCH_CHECK(s.device() == x.device(), "signs must reside on the same device as x");
|
||||
TORCH_CHECK(s.dim() == 4, "signs must be rank 4");
|
||||
TORCH_CHECK(s.size(0) == x.size(0) && s.size(1) == x.size(1), "signs must have same batch & channels as x");
|
||||
TORCH_CHECK(s.size(2) <= INT_MAX && (s.size(3) << 2) <= INT_MAX, "signs tensor is too large");
|
||||
}
|
||||
|
||||
// Initialize CUDA kernel parameters.
|
||||
filtered_lrelu_act_kernel_params p;
|
||||
p.x = x.data_ptr();
|
||||
p.s = (readSigns || writeSigns) ? s.data_ptr<unsigned char>() : 0;
|
||||
p.gain = gain;
|
||||
p.slope = slope;
|
||||
p.clamp = clamp;
|
||||
p.xShape = make_int4((int)x.size(3), (int)x.size(2), (int)x.size(1), (int)x.size(0));
|
||||
p.xStride = make_longlong4(x.stride(3), x.stride(2), x.stride(1), x.stride(0));
|
||||
p.sShape = (readSigns || writeSigns) ? make_int2((int)s.size(3) << 2, (int)s.size(2)) : make_int2(0, 0); // Width is in elements. Contiguous.
|
||||
p.sOfs = make_int2(sx, sy);
|
||||
|
||||
// Choose CUDA kernel.
|
||||
void* func = 0;
|
||||
AT_DISPATCH_FLOATING_TYPES_AND_HALF(x.scalar_type(), "filtered_lrelu_act_cuda", [&]
|
||||
{
|
||||
if (writeSigns)
|
||||
func = choose_filtered_lrelu_act_kernel<scalar_t, true, false>();
|
||||
else if (readSigns)
|
||||
func = choose_filtered_lrelu_act_kernel<scalar_t, false, true>();
|
||||
else
|
||||
func = choose_filtered_lrelu_act_kernel<scalar_t, false, false>();
|
||||
});
|
||||
TORCH_CHECK(func, "internal error - CUDA kernel not found");
|
||||
|
||||
// Launch CUDA kernel.
|
||||
void* args[] = {&p};
|
||||
int bx = 128; // 4 warps per block.
|
||||
|
||||
// Logical size of launch = writeSigns ? p.s : p.x
|
||||
uint32_t gx = writeSigns ? p.sShape.x : p.xShape.x;
|
||||
uint32_t gy = writeSigns ? p.sShape.y : p.xShape.y;
|
||||
uint32_t gz = p.xShape.z * p.xShape.w; // Same as in p.sShape if signs are in use.
|
||||
gx = (gx - 1) / bx + 1;
|
||||
|
||||
// Make sure grid y and z dimensions are within CUDA launch limits. Kernel loops internally to do the rest.
|
||||
const uint32_t gmax = 65535;
|
||||
gy = std::min(gy, gmax);
|
||||
gz = std::min(gz, gmax);
|
||||
|
||||
// Launch.
|
||||
AT_CUDA_CHECK(cudaLaunchKernel(func, dim3(gx, gy, gz), bx, args, 0, at::cuda::getCurrentCUDAStream()));
|
||||
return so;
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
|
||||
{
|
||||
m.def("filtered_lrelu", &filtered_lrelu); // The whole thing.
|
||||
m.def("filtered_lrelu_act_", &filtered_lrelu_act); // Activation and sign tensor handling only. Modifies data tensor in-place.
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,94 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
*
|
||||
* NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
* property and proprietary rights in and to this material, related
|
||||
* documentation and any modifications thereto. Any use, reproduction,
|
||||
* disclosure or distribution of this material and related documentation
|
||||
* without an express license agreement from NVIDIA CORPORATION or
|
||||
* its affiliates is strictly prohibited.
|
||||
*/
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// CUDA kernel parameters.
|
||||
|
||||
struct filtered_lrelu_kernel_params
|
||||
{
|
||||
// These parameters decide which kernel to use.
|
||||
int up; // upsampling ratio (1, 2, 4)
|
||||
int down; // downsampling ratio (1, 2, 4)
|
||||
int2 fuShape; // [size, 1] | [size, size]
|
||||
int2 fdShape; // [size, 1] | [size, size]
|
||||
|
||||
int _dummy; // Alignment.
|
||||
|
||||
// Rest of the parameters.
|
||||
const void* x; // Input tensor.
|
||||
void* y; // Output tensor.
|
||||
const void* b; // Bias tensor.
|
||||
unsigned char* s; // Sign tensor in/out. NULL if unused.
|
||||
const float* fu; // Upsampling filter.
|
||||
const float* fd; // Downsampling filter.
|
||||
|
||||
int2 pad0; // Left/top padding.
|
||||
float gain; // Additional gain factor.
|
||||
float slope; // Leaky ReLU slope on negative side.
|
||||
float clamp; // Clamp after nonlinearity.
|
||||
int flip; // Filter kernel flip for gradient computation.
|
||||
|
||||
int tilesXdim; // Original number of horizontal output tiles.
|
||||
int tilesXrep; // Number of horizontal tiles per CTA.
|
||||
int blockZofs; // Block z offset to support large minibatch, channel dimensions.
|
||||
|
||||
int4 xShape; // [width, height, channel, batch]
|
||||
int4 yShape; // [width, height, channel, batch]
|
||||
int2 sShape; // [width, height] - width is in bytes. Contiguous. Zeros if unused.
|
||||
int2 sOfs; // [ofs_x, ofs_y] - offset between upsampled data and sign tensor.
|
||||
int swLimit; // Active width of sign tensor in bytes.
|
||||
|
||||
longlong4 xStride; // Strides of all tensors except signs, same component order as shapes.
|
||||
longlong4 yStride; //
|
||||
int64_t bStride; //
|
||||
longlong3 fuStride; //
|
||||
longlong3 fdStride; //
|
||||
};
|
||||
|
||||
struct filtered_lrelu_act_kernel_params
|
||||
{
|
||||
void* x; // Input/output, modified in-place.
|
||||
unsigned char* s; // Sign tensor in/out. NULL if unused.
|
||||
|
||||
float gain; // Additional gain factor.
|
||||
float slope; // Leaky ReLU slope on negative side.
|
||||
float clamp; // Clamp after nonlinearity.
|
||||
|
||||
int4 xShape; // [width, height, channel, batch]
|
||||
longlong4 xStride; // Input/output tensor strides, same order as in shape.
|
||||
int2 sShape; // [width, height] - width is in elements. Contiguous. Zeros if unused.
|
||||
int2 sOfs; // [ofs_x, ofs_y] - offset between upsampled data and sign tensor.
|
||||
};
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// CUDA kernel specialization.
|
||||
|
||||
struct filtered_lrelu_kernel_spec
|
||||
{
|
||||
void* setup; // Function for filter kernel setup.
|
||||
void* exec; // Function for main operation.
|
||||
int2 tileOut; // Width/height of launch tile.
|
||||
int numWarps; // Number of warps per thread block, determines launch block size.
|
||||
int xrep; // For processing multiple horizontal tiles per thread block.
|
||||
int dynamicSharedKB; // How much dynamic shared memory the exec kernel wants.
|
||||
};
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// CUDA kernel selection.
|
||||
|
||||
template <class T, class index_t, bool signWrite, bool signRead> filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
template <class T, bool signWrite, bool signRead> void* choose_filtered_lrelu_act_kernel(void);
|
||||
template <bool signWrite, bool signRead> cudaError_t copy_filters(cudaStream_t stream);
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
@@ -0,0 +1,363 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
|
||||
import os
|
||||
import warnings
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from .. import custom_ops, misc
|
||||
from . import bias_act, upfirdn2d
|
||||
|
||||
_plugin = None
|
||||
|
||||
|
||||
def _init():
|
||||
global _plugin
|
||||
if _plugin is None:
|
||||
_plugin = custom_ops.get_plugin(
|
||||
module_name='filtered_lrelu_plugin',
|
||||
sources=[
|
||||
'filtered_lrelu.cpp', 'filtered_lrelu_wr.cu',
|
||||
'filtered_lrelu_rd.cu', 'filtered_lrelu_ns.cu'
|
||||
],
|
||||
headers=['filtered_lrelu.h', 'filtered_lrelu.cu'],
|
||||
source_dir=os.path.dirname(__file__),
|
||||
extra_cuda_cflags=['--use_fast_math'],
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def _get_filter_size(f):
|
||||
if f is None:
|
||||
return 1, 1
|
||||
assert isinstance(f, torch.Tensor)
|
||||
assert 1 <= f.ndim <= 2
|
||||
return f.shape[-1], f.shape[0] # width, height
|
||||
|
||||
|
||||
def _parse_padding(padding):
|
||||
if isinstance(padding, int):
|
||||
padding = [padding, padding]
|
||||
assert isinstance(padding, (list, tuple))
|
||||
assert all(isinstance(x, (int, np.integer)) for x in padding)
|
||||
padding = [int(x) for x in padding]
|
||||
if len(padding) == 2:
|
||||
px, py = padding
|
||||
padding = [px, px, py, py]
|
||||
px0, px1, py0, py1 = padding
|
||||
return px0, px1, py0, py1
|
||||
|
||||
|
||||
def filtered_lrelu(x,
|
||||
fu=None,
|
||||
fd=None,
|
||||
b=None,
|
||||
up=1,
|
||||
down=1,
|
||||
padding=0,
|
||||
gain=np.sqrt(2),
|
||||
slope=0.2,
|
||||
clamp=None,
|
||||
flip_filter=False,
|
||||
impl='cuda'):
|
||||
r"""Filtered leaky ReLU for a batch of 2D images.
|
||||
|
||||
Performs the following sequence of operations for each channel:
|
||||
|
||||
1. Add channel-specific bias if provided (`b`).
|
||||
|
||||
2. Upsample the image by inserting N-1 zeros after each pixel (`up`).
|
||||
|
||||
3. Pad the image with the specified number of zeros on each side (`padding`).
|
||||
Negative padding corresponds to cropping the image.
|
||||
|
||||
4. Convolve the image with the specified upsampling FIR filter (`fu`), shrinking it
|
||||
so that the footprint of all output pixels lies within the input image.
|
||||
|
||||
5. Multiply each value by the provided gain factor (`gain`).
|
||||
|
||||
6. Apply leaky ReLU activation function to each value.
|
||||
|
||||
7. Clamp each value between -clamp and +clamp, if `clamp` parameter is provided.
|
||||
|
||||
8. Convolve the image with the specified downsampling FIR filter (`fd`), shrinking
|
||||
it so that the footprint of all output pixels lies within the input image.
|
||||
|
||||
9. Downsample the image by keeping every Nth pixel (`down`).
|
||||
|
||||
The fused op is considerably more efficient than performing the same calculation
|
||||
using standard PyTorch ops. It supports gradients of arbitrary order.
|
||||
|
||||
Args:
|
||||
x: Float32/float16/float64 input tensor of the shape
|
||||
`[batch_size, num_channels, in_height, in_width]`.
|
||||
fu: Float32 upsampling FIR filter of the shape
|
||||
`[filter_height, filter_width]` (non-separable),
|
||||
`[filter_taps]` (separable), or
|
||||
`None` (identity).
|
||||
fd: Float32 downsampling FIR filter of the shape
|
||||
`[filter_height, filter_width]` (non-separable),
|
||||
`[filter_taps]` (separable), or
|
||||
`None` (identity).
|
||||
b: Bias vector, or `None` to disable. Must be a 1D tensor of the same type
|
||||
as `x`. The length of vector must must match the channel dimension of `x`.
|
||||
up: Integer upsampling factor (default: 1).
|
||||
down: Integer downsampling factor. (default: 1).
|
||||
padding: Padding with respect to the upsampled image. Can be a single number
|
||||
or a list/tuple `[x, y]` or `[x_before, x_after, y_before, y_after]`
|
||||
(default: 0).
|
||||
gain: Overall scaling factor for signal magnitude (default: sqrt(2)).
|
||||
slope: Slope on the negative side of leaky ReLU (default: 0.2).
|
||||
clamp: Maximum magnitude for leaky ReLU output (default: None).
|
||||
flip_filter: False = convolution, True = correlation (default: False).
|
||||
impl: Implementation to use. Can be `'ref'` or `'cuda'` (default: `'cuda'`).
|
||||
|
||||
Returns:
|
||||
Tensor of the shape `[batch_size, num_channels, out_height, out_width]`.
|
||||
"""
|
||||
assert isinstance(x, torch.Tensor)
|
||||
assert impl in ['ref', 'cuda']
|
||||
if impl == 'cuda' and x.device.type == 'cuda' and _init():
|
||||
return _filtered_lrelu_cuda(
|
||||
up=up,
|
||||
down=down,
|
||||
padding=padding,
|
||||
gain=gain,
|
||||
slope=slope,
|
||||
clamp=clamp,
|
||||
flip_filter=flip_filter).apply(x, fu, fd, b, None, 0, 0)
|
||||
return _filtered_lrelu_ref(
|
||||
x,
|
||||
fu=fu,
|
||||
fd=fd,
|
||||
b=b,
|
||||
up=up,
|
||||
down=down,
|
||||
padding=padding,
|
||||
gain=gain,
|
||||
slope=slope,
|
||||
clamp=clamp,
|
||||
flip_filter=flip_filter)
|
||||
|
||||
|
||||
@misc.profiled_function
|
||||
def _filtered_lrelu_ref(x,
|
||||
fu=None,
|
||||
fd=None,
|
||||
b=None,
|
||||
up=1,
|
||||
down=1,
|
||||
padding=0,
|
||||
gain=np.sqrt(2),
|
||||
slope=0.2,
|
||||
clamp=None,
|
||||
flip_filter=False):
|
||||
"""Slow and memory-inefficient reference implementation of `filtered_lrelu()` using
|
||||
existing `upfirdn2n()` and `bias_act()` ops.
|
||||
"""
|
||||
assert isinstance(x, torch.Tensor) and x.ndim == 4
|
||||
fu_w, fu_h = _get_filter_size(fu)
|
||||
fd_w, fd_h = _get_filter_size(fd)
|
||||
if b is not None:
|
||||
assert isinstance(b, torch.Tensor) and b.dtype == x.dtype
|
||||
misc.assert_shape(b, [x.shape[1]])
|
||||
assert isinstance(up, int) and up >= 1
|
||||
assert isinstance(down, int) and down >= 1
|
||||
px0, px1, py0, py1 = _parse_padding(padding)
|
||||
assert gain == float(gain) and gain > 0
|
||||
assert slope == float(slope) and slope >= 0
|
||||
assert clamp is None or (clamp == float(clamp) and clamp >= 0)
|
||||
|
||||
# Calculate output size.
|
||||
batch_size, channels, in_h, in_w = x.shape
|
||||
in_dtype = x.dtype
|
||||
temp_w = in_w * up + (px0 + px1) - (fu_w - 1) - (fd_w - 1) + (down - 1)
|
||||
out_w = temp_w // down
|
||||
temp_h = in_h * up + (py0 + py1) - (fu_h - 1) - (fd_h - 1) + (down - 1)
|
||||
out_h = temp_h // down
|
||||
|
||||
# Compute using existing ops.
|
||||
x = bias_act.bias_act(x=x, b=b) # Apply bias.
|
||||
x = upfirdn2d.upfirdn2d(
|
||||
x=x,
|
||||
f=fu,
|
||||
up=up,
|
||||
padding=[px0, px1, py0, py1],
|
||||
gain=up**2,
|
||||
flip_filter=flip_filter) # Upsample.
|
||||
x = bias_act.bias_act(
|
||||
x=x, act='lrelu', alpha=slope, gain=gain,
|
||||
clamp=clamp) # Bias, leaky ReLU, clamp.
|
||||
x = upfirdn2d.upfirdn2d(
|
||||
x=x, f=fd, down=down, flip_filter=flip_filter) # Downsample.
|
||||
|
||||
# Check output shape & dtype.
|
||||
misc.assert_shape(x, [batch_size, channels, out_h, out_w])
|
||||
assert x.dtype == in_dtype
|
||||
return x
|
||||
|
||||
|
||||
_filtered_lrelu_cuda_cache = dict()
|
||||
|
||||
|
||||
def _filtered_lrelu_cuda(up=1,
|
||||
down=1,
|
||||
padding=0,
|
||||
gain=np.sqrt(2),
|
||||
slope=0.2,
|
||||
clamp=None,
|
||||
flip_filter=False):
|
||||
"""Fast CUDA implementation of `filtered_lrelu()` using custom ops.
|
||||
"""
|
||||
assert isinstance(up, int) and up >= 1
|
||||
assert isinstance(down, int) and down >= 1
|
||||
px0, px1, py0, py1 = _parse_padding(padding)
|
||||
assert gain == float(gain) and gain > 0
|
||||
gain = float(gain)
|
||||
assert slope == float(slope) and slope >= 0
|
||||
slope = float(slope)
|
||||
assert clamp is None or (clamp == float(clamp) and clamp >= 0)
|
||||
clamp = float(clamp if clamp is not None else 'inf')
|
||||
|
||||
# Lookup from cache.
|
||||
key = (up, down, px0, px1, py0, py1, gain, slope, clamp, flip_filter)
|
||||
if key in _filtered_lrelu_cuda_cache:
|
||||
return _filtered_lrelu_cuda_cache[key]
|
||||
|
||||
# Forward op.
|
||||
class FilteredLReluCuda(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, x, fu, fd, b, si, sx, sy): # pylint: disable=arguments-differ
|
||||
assert isinstance(x, torch.Tensor) and x.ndim == 4
|
||||
|
||||
# Replace empty up/downsample kernels with full 1x1 kernels (faster than separable).
|
||||
if fu is None:
|
||||
fu = torch.ones([1, 1], dtype=torch.float32, device=x.device)
|
||||
if fd is None:
|
||||
fd = torch.ones([1, 1], dtype=torch.float32, device=x.device)
|
||||
assert 1 <= fu.ndim <= 2
|
||||
assert 1 <= fd.ndim <= 2
|
||||
|
||||
# Replace separable 1x1 kernels with full 1x1 kernels when scale factor is 1.
|
||||
if up == 1 and fu.ndim == 1 and fu.shape[0] == 1:
|
||||
fu = fu.square()[None]
|
||||
if down == 1 and fd.ndim == 1 and fd.shape[0] == 1:
|
||||
fd = fd.square()[None]
|
||||
|
||||
# Missing sign input tensor.
|
||||
if si is None:
|
||||
si = torch.empty([0])
|
||||
|
||||
# Missing bias tensor.
|
||||
if b is None:
|
||||
b = torch.zeros([x.shape[1]], dtype=x.dtype, device=x.device)
|
||||
|
||||
# Construct internal sign tensor only if gradients are needed.
|
||||
write_signs = (si.numel() == 0) and (x.requires_grad
|
||||
or b.requires_grad)
|
||||
|
||||
# Warn if input storage strides are not in decreasing order due to e.g. channels-last layout.
|
||||
strides = [x.stride(i) for i in range(x.ndim) if x.size(i) > 1]
|
||||
if any(a < b for a, b in zip(strides[:-1], strides[1:])):
|
||||
warnings.warn(
|
||||
'low-performance memory layout detected in filtered_lrelu input',
|
||||
RuntimeWarning)
|
||||
|
||||
# Call C++/Cuda plugin if datatype is supported.
|
||||
if x.dtype in [torch.float16, torch.float32]:
|
||||
if torch.cuda.current_stream(
|
||||
x.device) != torch.cuda.default_stream(x.device):
|
||||
warnings.warn(
|
||||
'filtered_lrelu called with non-default cuda stream but concurrent execution is not supported',
|
||||
RuntimeWarning)
|
||||
y, so, return_code = _plugin.filtered_lrelu(
|
||||
x, fu, fd, b, si, up, down, px0, px1, py0, py1, sx, sy,
|
||||
gain, slope, clamp, flip_filter, write_signs)
|
||||
else:
|
||||
return_code = -1
|
||||
|
||||
# only the bit-packed sign tensor is retained for gradient computation.
|
||||
if return_code < 0:
|
||||
warnings.warn(
|
||||
'filtered_lrelu called with parameters that have no optimized CUDA kernel, using generic fallback',
|
||||
RuntimeWarning)
|
||||
|
||||
y = x.add(b.unsqueeze(-1).unsqueeze(-1)) # Add bias.
|
||||
y = upfirdn2d.upfirdn2d(
|
||||
x=y,
|
||||
f=fu,
|
||||
up=up,
|
||||
padding=[px0, px1, py0, py1],
|
||||
gain=up**2,
|
||||
flip_filter=flip_filter) # Upsample.
|
||||
so = _plugin.filtered_lrelu_act_(
|
||||
y, si, sx, sy, gain, slope, clamp, write_signs
|
||||
) # Activation function and sign handling. Modifies y in-place.
|
||||
y = upfirdn2d.upfirdn2d(
|
||||
x=y, f=fd, down=down,
|
||||
flip_filter=flip_filter) # Downsample.
|
||||
|
||||
# Prepare for gradient computation.
|
||||
ctx.save_for_backward(fu, fd, (si if si.numel() else so))
|
||||
ctx.x_shape = x.shape
|
||||
ctx.y_shape = y.shape
|
||||
ctx.s_ofs = sx, sy
|
||||
return y
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, dy): # pylint: disable=arguments-differ
|
||||
fu, fd, si = ctx.saved_tensors
|
||||
_, _, xh, xw = ctx.x_shape
|
||||
_, _, yh, yw = ctx.y_shape
|
||||
sx, sy = ctx.s_ofs
|
||||
dx = None # 0
|
||||
dfu = None
|
||||
assert not ctx.needs_input_grad[1]
|
||||
dfd = None
|
||||
assert not ctx.needs_input_grad[2]
|
||||
db = None # 3
|
||||
dsi = None
|
||||
assert not ctx.needs_input_grad[4]
|
||||
dsx = None
|
||||
assert not ctx.needs_input_grad[5]
|
||||
dsy = None
|
||||
assert not ctx.needs_input_grad[6]
|
||||
|
||||
if ctx.needs_input_grad[0] or ctx.needs_input_grad[3]:
|
||||
pp = [
|
||||
(fu.shape[-1] - 1) + (fd.shape[-1] - 1) - px0,
|
||||
xw * up - yw * down + px0 - (up - 1),
|
||||
(fu.shape[0] - 1) + (fd.shape[0] - 1) - py0,
|
||||
xh * up - yh * down + py0 - (up - 1),
|
||||
]
|
||||
gg = gain * (up**2) / (down**2)
|
||||
ff = (not flip_filter)
|
||||
sx = sx - (fu.shape[-1] - 1) + px0
|
||||
sy = sy - (fu.shape[0] - 1) + py0
|
||||
dx = _filtered_lrelu_cuda(
|
||||
up=down,
|
||||
down=up,
|
||||
padding=pp,
|
||||
gain=gg,
|
||||
slope=slope,
|
||||
clamp=None,
|
||||
flip_filter=ff).apply(dy, fd, fu, None, si, sx, sy)
|
||||
|
||||
if ctx.needs_input_grad[3]:
|
||||
db = dx.sum([0, 2, 3])
|
||||
|
||||
return dx, dfu, dfd, db, dsi, dsx, dsy
|
||||
|
||||
# Add to cache.
|
||||
_filtered_lrelu_cuda_cache[key] = FilteredLReluCuda
|
||||
return FilteredLReluCuda
|
||||
@@ -0,0 +1,31 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
*
|
||||
* NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
* property and proprietary rights in and to this material, related
|
||||
* documentation and any modifications thereto. Any use, reproduction,
|
||||
* disclosure or distribution of this material and related documentation
|
||||
* without an express license agreement from NVIDIA CORPORATION or
|
||||
* its affiliates is strictly prohibited.
|
||||
*/
|
||||
|
||||
#include "filtered_lrelu.cu"
|
||||
|
||||
// Template/kernel specializations for no signs mode (no gradients required).
|
||||
|
||||
// Full op, 32-bit indexing.
|
||||
template filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel<c10::Half, int32_t, false, false>(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
template filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel<float, int32_t, false, false>(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
|
||||
// Full op, 64-bit indexing.
|
||||
template filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel<c10::Half, int64_t, false, false>(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
template filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel<float, int64_t, false, false>(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
|
||||
// Activation/signs only for generic variant. 64-bit indexing.
|
||||
template void* choose_filtered_lrelu_act_kernel<c10::Half, false, false>(void);
|
||||
template void* choose_filtered_lrelu_act_kernel<float, false, false>(void);
|
||||
template void* choose_filtered_lrelu_act_kernel<double, false, false>(void);
|
||||
|
||||
// Copy filters to constant memory.
|
||||
template cudaError_t copy_filters<false, false>(cudaStream_t stream);
|
||||
@@ -0,0 +1,31 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
*
|
||||
* NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
* property and proprietary rights in and to this material, related
|
||||
* documentation and any modifications thereto. Any use, reproduction,
|
||||
* disclosure or distribution of this material and related documentation
|
||||
* without an express license agreement from NVIDIA CORPORATION or
|
||||
* its affiliates is strictly prohibited.
|
||||
*/
|
||||
|
||||
#include "filtered_lrelu.cu"
|
||||
|
||||
// Template/kernel specializations for sign read mode.
|
||||
|
||||
// Full op, 32-bit indexing.
|
||||
template filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel<c10::Half, int32_t, false, true>(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
template filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel<float, int32_t, false, true>(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
|
||||
// Full op, 64-bit indexing.
|
||||
template filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel<c10::Half, int64_t, false, true>(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
template filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel<float, int64_t, false, true>(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
|
||||
// Activation/signs only for generic variant. 64-bit indexing.
|
||||
template void* choose_filtered_lrelu_act_kernel<c10::Half, false, true>(void);
|
||||
template void* choose_filtered_lrelu_act_kernel<float, false, true>(void);
|
||||
template void* choose_filtered_lrelu_act_kernel<double, false, true>(void);
|
||||
|
||||
// Copy filters to constant memory.
|
||||
template cudaError_t copy_filters<false, true>(cudaStream_t stream);
|
||||
@@ -0,0 +1,31 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
*
|
||||
* NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
* property and proprietary rights in and to this material, related
|
||||
* documentation and any modifications thereto. Any use, reproduction,
|
||||
* disclosure or distribution of this material and related documentation
|
||||
* without an express license agreement from NVIDIA CORPORATION or
|
||||
* its affiliates is strictly prohibited.
|
||||
*/
|
||||
|
||||
#include "filtered_lrelu.cu"
|
||||
|
||||
// Template/kernel specializations for sign write mode.
|
||||
|
||||
// Full op, 32-bit indexing.
|
||||
template filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel<c10::Half, int32_t, true, false>(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
template filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel<float, int32_t, true, false>(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
|
||||
// Full op, 64-bit indexing.
|
||||
template filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel<c10::Half, int64_t, true, false>(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
template filtered_lrelu_kernel_spec choose_filtered_lrelu_kernel<float, int64_t, true, false>(const filtered_lrelu_kernel_params& p, int sharedKB);
|
||||
|
||||
// Activation/signs only for generic variant. 64-bit indexing.
|
||||
template void* choose_filtered_lrelu_act_kernel<c10::Half, true, false>(void);
|
||||
template void* choose_filtered_lrelu_act_kernel<float, true, false>(void);
|
||||
template void* choose_filtered_lrelu_act_kernel<double, true, false>(void);
|
||||
|
||||
// Copy filters to constant memory.
|
||||
template cudaError_t copy_filters<true, false>(cudaStream_t stream);
|
||||
@@ -0,0 +1,60 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""Fused multiply-add, with slightly faster gradients than `torch.addcmul()`."""
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def fma(a, b, c): # => a * b + c
|
||||
return _FusedMultiplyAdd.apply(a, b, c)
|
||||
|
||||
|
||||
class _FusedMultiplyAdd(torch.autograd.Function): # a * b + c
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, a, b, c): # pylint: disable=arguments-differ
|
||||
out = torch.addcmul(c, a, b)
|
||||
ctx.save_for_backward(a, b)
|
||||
ctx.c_shape = c.shape
|
||||
return out
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, dout): # pylint: disable=arguments-differ
|
||||
a, b = ctx.saved_tensors
|
||||
c_shape = ctx.c_shape
|
||||
da = None
|
||||
db = None
|
||||
dc = None
|
||||
|
||||
if ctx.needs_input_grad[0]:
|
||||
da = _unbroadcast(dout * b, a.shape)
|
||||
|
||||
if ctx.needs_input_grad[1]:
|
||||
db = _unbroadcast(dout * a, b.shape)
|
||||
|
||||
if ctx.needs_input_grad[2]:
|
||||
dc = _unbroadcast(dout, c_shape)
|
||||
|
||||
return da, db, dc
|
||||
|
||||
|
||||
def _unbroadcast(x, shape):
|
||||
extra_dims = x.ndim - len(shape)
|
||||
assert extra_dims >= 0
|
||||
dim = [
|
||||
i for i in range(x.ndim)
|
||||
if x.shape[i] > 1 and (i < extra_dims or shape[i - extra_dims] == 1)
|
||||
]
|
||||
if len(dim):
|
||||
x = x.sum(dim=dim, keepdim=True)
|
||||
if extra_dims:
|
||||
x = x.reshape(-1, *x.shape[extra_dims + 1:])
|
||||
assert x.shape == shape
|
||||
return x
|
||||
@@ -0,0 +1,84 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""Custom replacement for `torch.nn.functional.grid_sample` that
|
||||
supports arbitrarily high order gradients between the input and output.
|
||||
Only works on 2D images and assumes
|
||||
`mode='bilinear'`, `padding_mode='zeros'`, `align_corners=False`."""
|
||||
|
||||
import torch
|
||||
|
||||
# pylint: disable=redefined-builtin
|
||||
# pylint: disable=arguments-differ
|
||||
# pylint: disable=protected-access
|
||||
|
||||
enabled = False # Enable the custom op by setting this to true.
|
||||
|
||||
|
||||
def grid_sample(input, grid):
|
||||
if _should_use_custom_op():
|
||||
return _GridSample2dForward.apply(input, grid)
|
||||
return torch.nn.functional.grid_sample(
|
||||
input=input,
|
||||
grid=grid,
|
||||
mode='bilinear',
|
||||
padding_mode='zeros',
|
||||
align_corners=False)
|
||||
|
||||
|
||||
def _should_use_custom_op():
|
||||
return enabled
|
||||
|
||||
|
||||
class _GridSample2dForward(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input, grid):
|
||||
assert input.ndim == 4
|
||||
assert grid.ndim == 4
|
||||
output = torch.nn.functional.grid_sample(
|
||||
input=input,
|
||||
grid=grid,
|
||||
mode='bilinear',
|
||||
padding_mode='zeros',
|
||||
align_corners=False)
|
||||
ctx.save_for_backward(input, grid)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
input, grid = ctx.saved_tensors
|
||||
grad_input, grad_grid = _GridSample2dBackward.apply(
|
||||
grad_output, input, grid)
|
||||
return grad_input, grad_grid
|
||||
|
||||
|
||||
class _GridSample2dBackward(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, grad_output, input, grid):
|
||||
op = torch._C._jit_get_operation('aten::grid_sampler_2d_backward')
|
||||
grad_input, grad_grid = op(grad_output, input, grid, 0, 0, False)
|
||||
ctx.save_for_backward(grid)
|
||||
return grad_input, grad_grid
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad2_grad_input, grad2_grad_grid):
|
||||
_ = grad2_grad_grid # unused
|
||||
grid, = ctx.saved_tensors
|
||||
grad2_grad_output = None
|
||||
grad2_input = None
|
||||
grad2_grid = None
|
||||
|
||||
if ctx.needs_input_grad[0]:
|
||||
grad2_grad_output = _GridSample2dForward.apply(
|
||||
grad2_grad_input, grid)
|
||||
|
||||
assert not ctx.needs_input_grad[2]
|
||||
return grad2_grad_output, grad2_input, grad2_grid
|
||||
@@ -0,0 +1,111 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
*
|
||||
* NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
* property and proprietary rights in and to this material, related
|
||||
* documentation and any modifications thereto. Any use, reproduction,
|
||||
* disclosure or distribution of this material and related documentation
|
||||
* without an express license agreement from NVIDIA CORPORATION or
|
||||
* its affiliates is strictly prohibited.
|
||||
*/
|
||||
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include "upfirdn2d.h"
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
|
||||
static torch::Tensor upfirdn2d(torch::Tensor x, torch::Tensor f, int upx, int upy, int downx, int downy, int padx0, int padx1, int pady0, int pady1, bool flip, float gain)
|
||||
{
|
||||
// Validate arguments.
|
||||
TORCH_CHECK(x.is_cuda(), "x must reside on CUDA device");
|
||||
TORCH_CHECK(f.device() == x.device(), "f must reside on the same device as x");
|
||||
TORCH_CHECK(f.dtype() == torch::kFloat, "f must be float32");
|
||||
TORCH_CHECK(x.numel() <= INT_MAX, "x is too large");
|
||||
TORCH_CHECK(f.numel() <= INT_MAX, "f is too large");
|
||||
TORCH_CHECK(x.numel() > 0, "x has zero size");
|
||||
TORCH_CHECK(f.numel() > 0, "f has zero size");
|
||||
TORCH_CHECK(x.dim() == 4, "x must be rank 4");
|
||||
TORCH_CHECK(f.dim() == 2, "f must be rank 2");
|
||||
TORCH_CHECK((x.size(0)-1)*x.stride(0) + (x.size(1)-1)*x.stride(1) + (x.size(2)-1)*x.stride(2) + (x.size(3)-1)*x.stride(3) <= INT_MAX, "x memory footprint is too large");
|
||||
TORCH_CHECK(f.size(0) >= 1 && f.size(1) >= 1, "f must be at least 1x1");
|
||||
TORCH_CHECK(upx >= 1 && upy >= 1, "upsampling factor must be at least 1");
|
||||
TORCH_CHECK(downx >= 1 && downy >= 1, "downsampling factor must be at least 1");
|
||||
|
||||
// Create output tensor.
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(x));
|
||||
int outW = ((int)x.size(3) * upx + padx0 + padx1 - (int)f.size(1) + downx) / downx;
|
||||
int outH = ((int)x.size(2) * upy + pady0 + pady1 - (int)f.size(0) + downy) / downy;
|
||||
TORCH_CHECK(outW >= 1 && outH >= 1, "output must be at least 1x1");
|
||||
torch::Tensor y = torch::empty({x.size(0), x.size(1), outH, outW}, x.options(), x.suggest_memory_format());
|
||||
TORCH_CHECK(y.numel() <= INT_MAX, "output is too large");
|
||||
TORCH_CHECK((y.size(0)-1)*y.stride(0) + (y.size(1)-1)*y.stride(1) + (y.size(2)-1)*y.stride(2) + (y.size(3)-1)*y.stride(3) <= INT_MAX, "output memory footprint is too large");
|
||||
|
||||
// Initialize CUDA kernel parameters.
|
||||
upfirdn2d_kernel_params p;
|
||||
p.x = x.data_ptr();
|
||||
p.f = f.data_ptr<float>();
|
||||
p.y = y.data_ptr();
|
||||
p.up = make_int2(upx, upy);
|
||||
p.down = make_int2(downx, downy);
|
||||
p.pad0 = make_int2(padx0, pady0);
|
||||
p.flip = (flip) ? 1 : 0;
|
||||
p.gain = gain;
|
||||
p.inSize = make_int4((int)x.size(3), (int)x.size(2), (int)x.size(1), (int)x.size(0));
|
||||
p.inStride = make_int4((int)x.stride(3), (int)x.stride(2), (int)x.stride(1), (int)x.stride(0));
|
||||
p.filterSize = make_int2((int)f.size(1), (int)f.size(0));
|
||||
p.filterStride = make_int2((int)f.stride(1), (int)f.stride(0));
|
||||
p.outSize = make_int4((int)y.size(3), (int)y.size(2), (int)y.size(1), (int)y.size(0));
|
||||
p.outStride = make_int4((int)y.stride(3), (int)y.stride(2), (int)y.stride(1), (int)y.stride(0));
|
||||
p.sizeMajor = (p.inStride.z == 1) ? p.inSize.w : p.inSize.w * p.inSize.z;
|
||||
p.sizeMinor = (p.inStride.z == 1) ? p.inSize.z : 1;
|
||||
|
||||
// Choose CUDA kernel.
|
||||
upfirdn2d_kernel_spec spec;
|
||||
AT_DISPATCH_FLOATING_TYPES_AND_HALF(x.scalar_type(), "upfirdn2d_cuda", [&]
|
||||
{
|
||||
spec = choose_upfirdn2d_kernel<scalar_t>(p);
|
||||
});
|
||||
|
||||
// Set looping options.
|
||||
p.loopMajor = (p.sizeMajor - 1) / 16384 + 1;
|
||||
p.loopMinor = spec.loopMinor;
|
||||
p.loopX = spec.loopX;
|
||||
p.launchMinor = (p.sizeMinor - 1) / p.loopMinor + 1;
|
||||
p.launchMajor = (p.sizeMajor - 1) / p.loopMajor + 1;
|
||||
|
||||
// Compute grid size.
|
||||
dim3 blockSize, gridSize;
|
||||
if (spec.tileOutW < 0) // large
|
||||
{
|
||||
blockSize = dim3(4, 32, 1);
|
||||
gridSize = dim3(
|
||||
((p.outSize.y - 1) / blockSize.x + 1) * p.launchMinor,
|
||||
(p.outSize.x - 1) / (blockSize.y * p.loopX) + 1,
|
||||
p.launchMajor);
|
||||
}
|
||||
else // small
|
||||
{
|
||||
blockSize = dim3(256, 1, 1);
|
||||
gridSize = dim3(
|
||||
((p.outSize.y - 1) / spec.tileOutH + 1) * p.launchMinor,
|
||||
(p.outSize.x - 1) / (spec.tileOutW * p.loopX) + 1,
|
||||
p.launchMajor);
|
||||
}
|
||||
|
||||
// Launch CUDA kernel.
|
||||
void* args[] = {&p};
|
||||
AT_CUDA_CHECK(cudaLaunchKernel(spec.kernel, gridSize, blockSize, args, 0, at::cuda::getCurrentCUDAStream()));
|
||||
return y;
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
|
||||
{
|
||||
m.def("upfirdn2d", &upfirdn2d);
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
@@ -0,0 +1,388 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
*
|
||||
* NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
* property and proprietary rights in and to this material, related
|
||||
* documentation and any modifications thereto. Any use, reproduction,
|
||||
* disclosure or distribution of this material and related documentation
|
||||
* without an express license agreement from NVIDIA CORPORATION or
|
||||
* its affiliates is strictly prohibited.
|
||||
*/
|
||||
|
||||
#include <c10/util/Half.h>
|
||||
#include "upfirdn2d.h"
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// Helpers.
|
||||
|
||||
template <class T> struct InternalType;
|
||||
template <> struct InternalType<double> { typedef double scalar_t; };
|
||||
template <> struct InternalType<float> { typedef float scalar_t; };
|
||||
template <> struct InternalType<c10::Half> { typedef float scalar_t; };
|
||||
|
||||
static __device__ __forceinline__ int floor_div(int a, int b)
|
||||
{
|
||||
int t = 1 - a / b;
|
||||
return (a + t * b) / b - t;
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// Generic CUDA implementation for large filters.
|
||||
|
||||
template <class T> static __global__ void upfirdn2d_kernel_large(upfirdn2d_kernel_params p)
|
||||
{
|
||||
typedef typename InternalType<T>::scalar_t scalar_t;
|
||||
|
||||
// Calculate thread index.
|
||||
int minorBase = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
int outY = minorBase / p.launchMinor;
|
||||
minorBase -= outY * p.launchMinor;
|
||||
int outXBase = blockIdx.y * p.loopX * blockDim.y + threadIdx.y;
|
||||
int majorBase = blockIdx.z * p.loopMajor;
|
||||
if (outXBase >= p.outSize.x | outY >= p.outSize.y | majorBase >= p.sizeMajor)
|
||||
return;
|
||||
|
||||
// Setup Y receptive field.
|
||||
int midY = outY * p.down.y + p.up.y - 1 - p.pad0.y;
|
||||
int inY = min(max(floor_div(midY, p.up.y), 0), p.inSize.y);
|
||||
int h = min(max(floor_div(midY + p.filterSize.y, p.up.y), 0), p.inSize.y) - inY;
|
||||
int filterY = midY + p.filterSize.y - (inY + 1) * p.up.y;
|
||||
if (p.flip)
|
||||
filterY = p.filterSize.y - 1 - filterY;
|
||||
|
||||
// Loop over major, minor, and X.
|
||||
for (int majorIdx = 0, major = majorBase; majorIdx < p.loopMajor & major < p.sizeMajor; majorIdx++, major++)
|
||||
for (int minorIdx = 0, minor = minorBase; minorIdx < p.loopMinor & minor < p.sizeMinor; minorIdx++, minor += p.launchMinor)
|
||||
{
|
||||
int nc = major * p.sizeMinor + minor;
|
||||
int n = nc / p.inSize.z;
|
||||
int c = nc - n * p.inSize.z;
|
||||
for (int loopX = 0, outX = outXBase; loopX < p.loopX & outX < p.outSize.x; loopX++, outX += blockDim.y)
|
||||
{
|
||||
// Setup X receptive field.
|
||||
int midX = outX * p.down.x + p.up.x - 1 - p.pad0.x;
|
||||
int inX = min(max(floor_div(midX, p.up.x), 0), p.inSize.x);
|
||||
int w = min(max(floor_div(midX + p.filterSize.x, p.up.x), 0), p.inSize.x) - inX;
|
||||
int filterX = midX + p.filterSize.x - (inX + 1) * p.up.x;
|
||||
if (p.flip)
|
||||
filterX = p.filterSize.x - 1 - filterX;
|
||||
|
||||
// Initialize pointers.
|
||||
const T* xp = &((const T*)p.x)[inX * p.inStride.x + inY * p.inStride.y + c * p.inStride.z + n * p.inStride.w];
|
||||
const float* fp = &p.f[filterX * p.filterStride.x + filterY * p.filterStride.y];
|
||||
int filterStepX = ((p.flip) ? p.up.x : -p.up.x) * p.filterStride.x;
|
||||
int filterStepY = ((p.flip) ? p.up.y : -p.up.y) * p.filterStride.y;
|
||||
|
||||
// Inner loop.
|
||||
scalar_t v = 0;
|
||||
for (int y = 0; y < h; y++)
|
||||
{
|
||||
for (int x = 0; x < w; x++)
|
||||
{
|
||||
v += (scalar_t)(*xp) * (scalar_t)(*fp);
|
||||
xp += p.inStride.x;
|
||||
fp += filterStepX;
|
||||
}
|
||||
xp += p.inStride.y - w * p.inStride.x;
|
||||
fp += filterStepY - w * filterStepX;
|
||||
}
|
||||
|
||||
// Store result.
|
||||
v *= p.gain;
|
||||
((T*)p.y)[outX * p.outStride.x + outY * p.outStride.y + c * p.outStride.z + n * p.outStride.w] = (T)v;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// Specialized CUDA implementation for small filters.
|
||||
|
||||
template <class T, int upx, int upy, int downx, int downy, int filterW, int filterH, int tileOutW, int tileOutH, int loopMinor>
|
||||
static __global__ void upfirdn2d_kernel_small(upfirdn2d_kernel_params p)
|
||||
{
|
||||
typedef typename InternalType<T>::scalar_t scalar_t;
|
||||
const int tileInW = ((tileOutW - 1) * downx + filterW - 1) / upx + 1;
|
||||
const int tileInH = ((tileOutH - 1) * downy + filterH - 1) / upy + 1;
|
||||
__shared__ volatile scalar_t sf[filterH][filterW];
|
||||
__shared__ volatile scalar_t sx[tileInH][tileInW][loopMinor];
|
||||
|
||||
// Calculate tile index.
|
||||
int minorBase = blockIdx.x;
|
||||
int tileOutY = minorBase / p.launchMinor;
|
||||
minorBase -= tileOutY * p.launchMinor;
|
||||
minorBase *= loopMinor;
|
||||
tileOutY *= tileOutH;
|
||||
int tileOutXBase = blockIdx.y * p.loopX * tileOutW;
|
||||
int majorBase = blockIdx.z * p.loopMajor;
|
||||
if (tileOutXBase >= p.outSize.x | tileOutY >= p.outSize.y | majorBase >= p.sizeMajor)
|
||||
return;
|
||||
|
||||
// Load filter (flipped).
|
||||
for (int tapIdx = threadIdx.x; tapIdx < filterH * filterW; tapIdx += blockDim.x)
|
||||
{
|
||||
int fy = tapIdx / filterW;
|
||||
int fx = tapIdx - fy * filterW;
|
||||
scalar_t v = 0;
|
||||
if (fx < p.filterSize.x & fy < p.filterSize.y)
|
||||
{
|
||||
int ffx = (p.flip) ? fx : p.filterSize.x - 1 - fx;
|
||||
int ffy = (p.flip) ? fy : p.filterSize.y - 1 - fy;
|
||||
v = (scalar_t)p.f[ffx * p.filterStride.x + ffy * p.filterStride.y];
|
||||
}
|
||||
sf[fy][fx] = v;
|
||||
}
|
||||
|
||||
// Loop over major and X.
|
||||
for (int majorIdx = 0, major = majorBase; majorIdx < p.loopMajor & major < p.sizeMajor; majorIdx++, major++)
|
||||
{
|
||||
int baseNC = major * p.sizeMinor + minorBase;
|
||||
int n = baseNC / p.inSize.z;
|
||||
int baseC = baseNC - n * p.inSize.z;
|
||||
for (int loopX = 0, tileOutX = tileOutXBase; loopX < p.loopX & tileOutX < p.outSize.x; loopX++, tileOutX += tileOutW)
|
||||
{
|
||||
// Load input pixels.
|
||||
int tileMidX = tileOutX * downx + upx - 1 - p.pad0.x;
|
||||
int tileMidY = tileOutY * downy + upy - 1 - p.pad0.y;
|
||||
int tileInX = floor_div(tileMidX, upx);
|
||||
int tileInY = floor_div(tileMidY, upy);
|
||||
__syncthreads();
|
||||
for (int inIdx = threadIdx.x; inIdx < tileInH * tileInW * loopMinor; inIdx += blockDim.x)
|
||||
{
|
||||
int relC = inIdx;
|
||||
int relInX = relC / loopMinor;
|
||||
int relInY = relInX / tileInW;
|
||||
relC -= relInX * loopMinor;
|
||||
relInX -= relInY * tileInW;
|
||||
int c = baseC + relC;
|
||||
int inX = tileInX + relInX;
|
||||
int inY = tileInY + relInY;
|
||||
scalar_t v = 0;
|
||||
if (inX >= 0 & inY >= 0 & inX < p.inSize.x & inY < p.inSize.y & c < p.inSize.z)
|
||||
v = (scalar_t)((const T*)p.x)[inX * p.inStride.x + inY * p.inStride.y + c * p.inStride.z + n * p.inStride.w];
|
||||
sx[relInY][relInX][relC] = v;
|
||||
}
|
||||
|
||||
// Loop over output pixels.
|
||||
__syncthreads();
|
||||
for (int outIdx = threadIdx.x; outIdx < tileOutH * tileOutW * loopMinor; outIdx += blockDim.x)
|
||||
{
|
||||
int relC = outIdx;
|
||||
int relOutX = relC / loopMinor;
|
||||
int relOutY = relOutX / tileOutW;
|
||||
relC -= relOutX * loopMinor;
|
||||
relOutX -= relOutY * tileOutW;
|
||||
int c = baseC + relC;
|
||||
int outX = tileOutX + relOutX;
|
||||
int outY = tileOutY + relOutY;
|
||||
|
||||
// Setup receptive field.
|
||||
int midX = tileMidX + relOutX * downx;
|
||||
int midY = tileMidY + relOutY * downy;
|
||||
int inX = floor_div(midX, upx);
|
||||
int inY = floor_div(midY, upy);
|
||||
int relInX = inX - tileInX;
|
||||
int relInY = inY - tileInY;
|
||||
int filterX = (inX + 1) * upx - midX - 1; // flipped
|
||||
int filterY = (inY + 1) * upy - midY - 1; // flipped
|
||||
|
||||
// Inner loop.
|
||||
if (outX < p.outSize.x & outY < p.outSize.y & c < p.outSize.z)
|
||||
{
|
||||
scalar_t v = 0;
|
||||
#pragma unroll
|
||||
for (int y = 0; y < filterH / upy; y++)
|
||||
#pragma unroll
|
||||
for (int x = 0; x < filterW / upx; x++)
|
||||
v += sx[relInY + y][relInX + x][relC] * sf[filterY + y * upy][filterX + x * upx];
|
||||
v *= p.gain;
|
||||
((T*)p.y)[outX * p.outStride.x + outY * p.outStride.y + c * p.outStride.z + n * p.outStride.w] = (T)v;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// CUDA kernel selection.
|
||||
|
||||
template <class T> upfirdn2d_kernel_spec choose_upfirdn2d_kernel(const upfirdn2d_kernel_params& p)
|
||||
{
|
||||
int s = p.inStride.z, fx = p.filterSize.x, fy = p.filterSize.y;
|
||||
upfirdn2d_kernel_spec spec = {(void*)upfirdn2d_kernel_large<T>, -1,-1,1, 4}; // contiguous
|
||||
if (s == 1) spec = {(void*)upfirdn2d_kernel_large<T>, -1,-1,4, 1}; // channels_last
|
||||
|
||||
// No up/downsampling.
|
||||
if (p.up.x == 1 && p.up.y == 1 && p.down.x == 1 && p.down.y == 1)
|
||||
{
|
||||
// contiguous
|
||||
if (s != 1 && fx <= 24 && fy <= 24) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 24,24, 64,32,1>, 64,32,1, 1};
|
||||
if (s != 1 && fx <= 16 && fy <= 16) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 16,16, 64,32,1>, 64,32,1, 1};
|
||||
if (s != 1 && fx <= 7 && fy <= 7 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 7,7, 64,16,1>, 64,16,1, 1};
|
||||
if (s != 1 && fx <= 6 && fy <= 6 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 6,6, 64,16,1>, 64,16,1, 1};
|
||||
if (s != 1 && fx <= 5 && fy <= 5 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 5,5, 64,16,1>, 64,16,1, 1};
|
||||
if (s != 1 && fx <= 4 && fy <= 4 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 4,4, 64,16,1>, 64,16,1, 1};
|
||||
if (s != 1 && fx <= 3 && fy <= 3 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 3,3, 64,16,1>, 64,16,1, 1};
|
||||
if (s != 1 && fx <= 24 && fy <= 1 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 24,1, 128,8,1>, 128,8,1, 1};
|
||||
if (s != 1 && fx <= 16 && fy <= 1 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 16,1, 128,8,1>, 128,8,1, 1};
|
||||
if (s != 1 && fx <= 8 && fy <= 1 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 8,1, 128,8,1>, 128,8,1, 1};
|
||||
if (s != 1 && fx <= 1 && fy <= 24) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 1,24, 32,32,1>, 32,32,1, 1};
|
||||
if (s != 1 && fx <= 1 && fy <= 16) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 1,16, 32,32,1>, 32,32,1, 1};
|
||||
if (s != 1 && fx <= 1 && fy <= 8 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 1,8, 32,32,1>, 32,32,1, 1};
|
||||
// channels_last
|
||||
if (s == 1 && fx <= 24 && fy <= 24) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 24,24, 32,32,1>, 32,32,1, 1};
|
||||
if (s == 1 && fx <= 16 && fy <= 16) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 16,16, 32,32,1>, 32,32,1, 1};
|
||||
if (s == 1 && fx <= 7 && fy <= 7 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 7,7, 16,16,8>, 16,16,8, 1};
|
||||
if (s == 1 && fx <= 6 && fy <= 6 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 6,6, 16,16,8>, 16,16,8, 1};
|
||||
if (s == 1 && fx <= 5 && fy <= 5 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 5,5, 16,16,8>, 16,16,8, 1};
|
||||
if (s == 1 && fx <= 4 && fy <= 4 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 4,4, 16,16,8>, 16,16,8, 1};
|
||||
if (s == 1 && fx <= 3 && fy <= 3 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 3,3, 16,16,8>, 16,16,8, 1};
|
||||
if (s == 1 && fx <= 24 && fy <= 1 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 24,1, 128,1,16>, 128,1,16, 1};
|
||||
if (s == 1 && fx <= 16 && fy <= 1 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 16,1, 128,1,16>, 128,1,16, 1};
|
||||
if (s == 1 && fx <= 8 && fy <= 1 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 8,1, 128,1,16>, 128,1,16, 1};
|
||||
if (s == 1 && fx <= 1 && fy <= 24) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 1,24, 1,128,16>, 1,128,16, 1};
|
||||
if (s == 1 && fx <= 1 && fy <= 16) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 1,16, 1,128,16>, 1,128,16, 1};
|
||||
if (s == 1 && fx <= 1 && fy <= 8 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,1, 1,8, 1,128,16>, 1,128,16, 1};
|
||||
}
|
||||
|
||||
// 2x upsampling.
|
||||
if (p.up.x == 2 && p.up.y == 2 && p.down.x == 1 && p.down.y == 1)
|
||||
{
|
||||
// contiguous
|
||||
if (s != 1 && fx <= 24 && fy <= 24) spec = {(void*)upfirdn2d_kernel_small<T, 2,2, 1,1, 24,24, 64,32,1>, 64,32,1, 1};
|
||||
if (s != 1 && fx <= 16 && fy <= 16) spec = {(void*)upfirdn2d_kernel_small<T, 2,2, 1,1, 16,16, 64,32,1>, 64,32,1, 1};
|
||||
if (s != 1 && fx <= 8 && fy <= 8 ) spec = {(void*)upfirdn2d_kernel_small<T, 2,2, 1,1, 8,8, 64,16,1>, 64,16,1, 1};
|
||||
if (s != 1 && fx <= 6 && fy <= 6 ) spec = {(void*)upfirdn2d_kernel_small<T, 2,2, 1,1, 6,6, 64,16,1>, 64,16,1, 1};
|
||||
if (s != 1 && fx <= 4 && fy <= 4 ) spec = {(void*)upfirdn2d_kernel_small<T, 2,2, 1,1, 4,4, 64,16,1>, 64,16,1, 1};
|
||||
if (s != 1 && fx <= 2 && fy <= 2 ) spec = {(void*)upfirdn2d_kernel_small<T, 2,2, 1,1, 2,2, 64,16,1>, 64,16,1, 1};
|
||||
// channels_last
|
||||
if (s == 1 && fx <= 24 && fy <= 24) spec = {(void*)upfirdn2d_kernel_small<T, 2,2, 1,1, 24,24, 32,32,1>, 32,32,1, 1};
|
||||
if (s == 1 && fx <= 16 && fy <= 16) spec = {(void*)upfirdn2d_kernel_small<T, 2,2, 1,1, 16,16, 32,32,1>, 32,32,1, 1};
|
||||
if (s == 1 && fx <= 8 && fy <= 8 ) spec = {(void*)upfirdn2d_kernel_small<T, 2,2, 1,1, 8,8, 16,16,8>, 16,16,8, 1};
|
||||
if (s == 1 && fx <= 6 && fy <= 6 ) spec = {(void*)upfirdn2d_kernel_small<T, 2,2, 1,1, 6,6, 16,16,8>, 16,16,8, 1};
|
||||
if (s == 1 && fx <= 4 && fy <= 4 ) spec = {(void*)upfirdn2d_kernel_small<T, 2,2, 1,1, 4,4, 16,16,8>, 16,16,8, 1};
|
||||
if (s == 1 && fx <= 2 && fy <= 2 ) spec = {(void*)upfirdn2d_kernel_small<T, 2,2, 1,1, 2,2, 16,16,8>, 16,16,8, 1};
|
||||
}
|
||||
if (p.up.x == 2 && p.up.y == 1 && p.down.x == 1 && p.down.y == 1)
|
||||
{
|
||||
// contiguous
|
||||
if (s != 1 && fx <= 24 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 2,1, 1,1, 24,1, 128,8,1>, 128,8,1, 1};
|
||||
if (s != 1 && fx <= 16 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 2,1, 1,1, 16,1, 128,8,1>, 128,8,1, 1};
|
||||
if (s != 1 && fx <= 8 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 2,1, 1,1, 8,1, 128,8,1>, 128,8,1, 1};
|
||||
// channels_last
|
||||
if (s == 1 && fx <= 24 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 2,1, 1,1, 24,1, 128,1,16>, 128,1,16, 1};
|
||||
if (s == 1 && fx <= 16 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 2,1, 1,1, 16,1, 128,1,16>, 128,1,16, 1};
|
||||
if (s == 1 && fx <= 8 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 2,1, 1,1, 8,1, 128,1,16>, 128,1,16, 1};
|
||||
}
|
||||
if (p.up.x == 1 && p.up.y == 2 && p.down.x == 1 && p.down.y == 1)
|
||||
{
|
||||
// contiguous
|
||||
if (s != 1 && fx <= 1 && fy <= 24) spec = {(void*)upfirdn2d_kernel_small<T, 1,2, 1,1, 1,24, 32,32,1>, 32,32,1, 1};
|
||||
if (s != 1 && fx <= 1 && fy <= 16) spec = {(void*)upfirdn2d_kernel_small<T, 1,2, 1,1, 1,16, 32,32,1>, 32,32,1, 1};
|
||||
if (s != 1 && fx <= 1 && fy <= 8 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,2, 1,1, 1,8, 32,32,1>, 32,32,1, 1};
|
||||
// channels_last
|
||||
if (s == 1 && fx <= 1 && fy <= 24) spec = {(void*)upfirdn2d_kernel_small<T, 1,2, 1,1, 1,24, 1,128,16>, 1,128,16, 1};
|
||||
if (s == 1 && fx <= 1 && fy <= 16) spec = {(void*)upfirdn2d_kernel_small<T, 1,2, 1,1, 1,16, 1,128,16>, 1,128,16, 1};
|
||||
if (s == 1 && fx <= 1 && fy <= 8 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,2, 1,1, 1,8, 1,128,16>, 1,128,16, 1};
|
||||
}
|
||||
|
||||
// 2x downsampling.
|
||||
if (p.up.x == 1 && p.up.y == 1 && p.down.x == 2 && p.down.y == 2)
|
||||
{
|
||||
// contiguous
|
||||
if (s != 1 && fx <= 24 && fy <= 24) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,2, 24,24, 32,16,1>, 32,16,1, 1};
|
||||
if (s != 1 && fx <= 16 && fy <= 16) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,2, 16,16, 32,16,1>, 32,16,1, 1};
|
||||
if (s != 1 && fx <= 8 && fy <= 8 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,2, 8,8, 32,8,1>, 32,8,1, 1};
|
||||
if (s != 1 && fx <= 6 && fy <= 6 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,2, 6,6, 32,8,1>, 32,8,1, 1};
|
||||
if (s != 1 && fx <= 4 && fy <= 4 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,2, 4,4, 32,8,1>, 32,8,1, 1};
|
||||
if (s != 1 && fx <= 2 && fy <= 2 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,2, 2,2, 32,8,1>, 32,8,1, 1};
|
||||
// channels_last
|
||||
if (s == 1 && fx <= 24 && fy <= 24) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,2, 24,24, 16,16,1>, 16,16,1, 1};
|
||||
if (s == 1 && fx <= 16 && fy <= 16) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,2, 16,16, 16,16,1>, 16,16,1, 1};
|
||||
if (s == 1 && fx <= 8 && fy <= 8 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,2, 8,8, 8,8,8>, 8,8,8, 1};
|
||||
if (s == 1 && fx <= 6 && fy <= 6 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,2, 6,6, 8,8,8>, 8,8,8, 1};
|
||||
if (s == 1 && fx <= 4 && fy <= 4 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,2, 4,4, 8,8,8>, 8,8,8, 1};
|
||||
if (s == 1 && fx <= 2 && fy <= 2 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,2, 2,2, 8,8,8>, 8,8,8, 1};
|
||||
}
|
||||
if (p.up.x == 1 && p.up.y == 1 && p.down.x == 2 && p.down.y == 1)
|
||||
{
|
||||
// contiguous
|
||||
if (s != 1 && fx <= 24 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,1, 24,1, 64,8,1>, 64,8,1, 1};
|
||||
if (s != 1 && fx <= 16 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,1, 16,1, 64,8,1>, 64,8,1, 1};
|
||||
if (s != 1 && fx <= 8 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,1, 8,1, 64,8,1>, 64,8,1, 1};
|
||||
// channels_last
|
||||
if (s == 1 && fx <= 24 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,1, 24,1, 64,1,8>, 64,1,8, 1};
|
||||
if (s == 1 && fx <= 16 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,1, 16,1, 64,1,8>, 64,1,8, 1};
|
||||
if (s == 1 && fx <= 8 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 2,1, 8,1, 64,1,8>, 64,1,8, 1};
|
||||
}
|
||||
if (p.up.x == 1 && p.up.y == 1 && p.down.x == 1 && p.down.y == 2)
|
||||
{
|
||||
// contiguous
|
||||
if (s != 1 && fx <= 1 && fy <= 24) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,2, 1,24, 32,16,1>, 32,16,1, 1};
|
||||
if (s != 1 && fx <= 1 && fy <= 16) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,2, 1,16, 32,16,1>, 32,16,1, 1};
|
||||
if (s != 1 && fx <= 1 && fy <= 8 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,2, 1,8, 32,16,1>, 32,16,1, 1};
|
||||
// channels_last
|
||||
if (s == 1 && fx <= 1 && fy <= 24) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,2, 1,24, 1,64,8>, 1,64,8, 1};
|
||||
if (s == 1 && fx <= 1 && fy <= 16) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,2, 1,16, 1,64,8>, 1,64,8, 1};
|
||||
if (s == 1 && fx <= 1 && fy <= 8 ) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,2, 1,8, 1,64,8>, 1,64,8, 1};
|
||||
}
|
||||
|
||||
// 4x upsampling.
|
||||
if (p.up.x == 4 && p.up.y == 4 && p.down.x == 1 && p.down.y == 1)
|
||||
{
|
||||
// contiguous
|
||||
if (s != 1 && fx <= 48 && fy <= 48) spec = {(void*)upfirdn2d_kernel_small<T, 4,4, 1,1, 48,48, 64,32,1>, 64,32,1, 1};
|
||||
if (s != 1 && fx <= 32 && fy <= 32) spec = {(void*)upfirdn2d_kernel_small<T, 4,4, 1,1, 32,32, 64,32,1>, 64,32,1, 1};
|
||||
// channels_last
|
||||
if (s == 1 && fx <= 48 && fy <= 48) spec = {(void*)upfirdn2d_kernel_small<T, 4,4, 1,1, 48,48, 32,32,1>, 32,32,1, 1};
|
||||
if (s == 1 && fx <= 32 && fy <= 32) spec = {(void*)upfirdn2d_kernel_small<T, 4,4, 1,1, 32,32, 32,32,1>, 32,32,1, 1};
|
||||
}
|
||||
if (p.up.x == 4 && p.up.y == 1 && p.down.x == 1 && p.down.y == 1)
|
||||
{
|
||||
// contiguous
|
||||
if (s != 1 && fx <= 48 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 4,1, 1,1, 48,1, 128,8,1>, 128,8,1, 1};
|
||||
if (s != 1 && fx <= 32 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 4,1, 1,1, 32,1, 128,8,1>, 128,8,1, 1};
|
||||
// channels_last
|
||||
if (s == 1 && fx <= 48 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 4,1, 1,1, 48,1, 128,1,16>, 128,1,16, 1};
|
||||
if (s == 1 && fx <= 32 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 4,1, 1,1, 32,1, 128,1,16>, 128,1,16, 1};
|
||||
}
|
||||
if (p.up.x == 1 && p.up.y == 4 && p.down.x == 1 && p.down.y == 1)
|
||||
{
|
||||
// contiguous
|
||||
if (s != 1 && fx <= 1 && fy <= 48) spec = {(void*)upfirdn2d_kernel_small<T, 1,4, 1,1, 1,48, 32,32,1>, 32,32,1, 1};
|
||||
if (s != 1 && fx <= 1 && fy <= 32) spec = {(void*)upfirdn2d_kernel_small<T, 1,4, 1,1, 1,32, 32,32,1>, 32,32,1, 1};
|
||||
// channels_last
|
||||
if (s == 1 && fx <= 1 && fy <= 48) spec = {(void*)upfirdn2d_kernel_small<T, 1,4, 1,1, 1,48, 1,128,16>, 1,128,16, 1};
|
||||
if (s == 1 && fx <= 1 && fy <= 32) spec = {(void*)upfirdn2d_kernel_small<T, 1,4, 1,1, 1,32, 1,128,16>, 1,128,16, 1};
|
||||
}
|
||||
|
||||
// 4x downsampling (inefficient).
|
||||
if (p.up.x == 1 && p.up.y == 1 && p.down.x == 4 && p.down.y == 1)
|
||||
{
|
||||
// contiguous
|
||||
if (s != 1 && fx <= 48 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 4,1, 48,1, 32,8,1>, 32,8,1, 1};
|
||||
if (s != 1 && fx <= 32 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 4,1, 32,1, 32,8,1>, 32,8,1, 1};
|
||||
// channels_last
|
||||
if (s == 1 && fx <= 48 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 4,1, 48,1, 32,1,8>, 32,1,8, 1};
|
||||
if (s == 1 && fx <= 32 && fy <= 1) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 4,1, 32,1, 32,1,8>, 32,1,8, 1};
|
||||
}
|
||||
if (p.up.x == 1 && p.up.y == 1 && p.down.x == 1 && p.down.y == 4)
|
||||
{
|
||||
// contiguous
|
||||
if (s != 1 && fx <= 1 && fy <= 48) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,4, 1,48, 32,8,1>, 32,8,1, 1};
|
||||
if (s != 1 && fx <= 1 && fy <= 32) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,4, 1,32, 32,8,1>, 32,8,1, 1};
|
||||
// channels_last
|
||||
if (s == 1 && fx <= 1 && fy <= 48) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,4, 1,48, 1,32,8>, 1,32,8, 1};
|
||||
if (s == 1 && fx <= 1 && fy <= 32) spec = {(void*)upfirdn2d_kernel_small<T, 1,1, 1,4, 1,32, 1,32,8>, 1,32,8, 1};
|
||||
}
|
||||
return spec;
|
||||
}
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// Template specializations.
|
||||
|
||||
template upfirdn2d_kernel_spec choose_upfirdn2d_kernel<double> (const upfirdn2d_kernel_params& p);
|
||||
template upfirdn2d_kernel_spec choose_upfirdn2d_kernel<float> (const upfirdn2d_kernel_params& p);
|
||||
template upfirdn2d_kernel_spec choose_upfirdn2d_kernel<c10::Half>(const upfirdn2d_kernel_params& p);
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
@@ -0,0 +1,63 @@
|
||||
/*
|
||||
* SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
*
|
||||
* NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
* property and proprietary rights in and to this material, related
|
||||
* documentation and any modifications thereto. Any use, reproduction,
|
||||
* disclosure or distribution of this material and related documentation
|
||||
* without an express license agreement from NVIDIA CORPORATION or
|
||||
* its affiliates is strictly prohibited.
|
||||
*/
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// CUDA kernel parameters.
|
||||
|
||||
struct upfirdn2d_kernel_params
|
||||
{
|
||||
const void* x;
|
||||
const float* f;
|
||||
void* y;
|
||||
|
||||
int2 up;
|
||||
int2 down;
|
||||
int2 pad0;
|
||||
int flip;
|
||||
float gain;
|
||||
|
||||
int4 inSize; // [width, height, channel, batch]
|
||||
int4 inStride;
|
||||
int2 filterSize; // [width, height]
|
||||
int2 filterStride;
|
||||
int4 outSize; // [width, height, channel, batch]
|
||||
int4 outStride;
|
||||
int sizeMinor;
|
||||
int sizeMajor;
|
||||
|
||||
int loopMinor;
|
||||
int loopMajor;
|
||||
int loopX;
|
||||
int launchMinor;
|
||||
int launchMajor;
|
||||
};
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// CUDA kernel specialization.
|
||||
|
||||
struct upfirdn2d_kernel_spec
|
||||
{
|
||||
void* kernel;
|
||||
int tileOutW;
|
||||
int tileOutH;
|
||||
int loopMinor;
|
||||
int loopX;
|
||||
};
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
// CUDA kernel selection.
|
||||
|
||||
template <class T> upfirdn2d_kernel_spec choose_upfirdn2d_kernel(const upfirdn2d_kernel_params& p);
|
||||
|
||||
//------------------------------------------------------------------------
|
||||
@@ -0,0 +1,448 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""Custom PyTorch ops for efficient resampling of 2D images."""
|
||||
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from .. import custom_ops, misc
|
||||
from . import conv2d_gradfix
|
||||
|
||||
_plugin = None
|
||||
|
||||
|
||||
def _init():
|
||||
global _plugin
|
||||
if _plugin is None:
|
||||
_plugin = custom_ops.get_plugin(
|
||||
module_name='upfirdn2d_plugin',
|
||||
sources=['upfirdn2d.cpp', 'upfirdn2d.cu'],
|
||||
headers=['upfirdn2d.h'],
|
||||
source_dir=os.path.dirname(__file__),
|
||||
extra_cuda_cflags=['--use_fast_math'],
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def _parse_scaling(scaling):
|
||||
if isinstance(scaling, int):
|
||||
scaling = [scaling, scaling]
|
||||
assert isinstance(scaling, (list, tuple))
|
||||
assert all(isinstance(x, int) for x in scaling)
|
||||
sx, sy = scaling
|
||||
assert sx >= 1 and sy >= 1
|
||||
return sx, sy
|
||||
|
||||
|
||||
def _parse_padding(padding):
|
||||
if isinstance(padding, int):
|
||||
padding = [padding, padding]
|
||||
assert isinstance(padding, (list, tuple))
|
||||
assert all(isinstance(x, int) for x in padding)
|
||||
if len(padding) == 2:
|
||||
padx, pady = padding
|
||||
padding = [padx, padx, pady, pady]
|
||||
padx0, padx1, pady0, pady1 = padding
|
||||
return padx0, padx1, pady0, pady1
|
||||
|
||||
|
||||
def _get_filter_size(f):
|
||||
if f is None:
|
||||
return 1, 1
|
||||
assert isinstance(f, torch.Tensor) and f.ndim in [1, 2]
|
||||
fw = f.shape[-1]
|
||||
fh = f.shape[0]
|
||||
with misc.suppress_tracer_warnings():
|
||||
fw = int(fw)
|
||||
fh = int(fh)
|
||||
misc.assert_shape(f, [fh, fw][:f.ndim])
|
||||
assert fw >= 1 and fh >= 1
|
||||
return fw, fh
|
||||
|
||||
|
||||
def setup_filter(f,
|
||||
device=torch.device('cpu'),
|
||||
normalize=True,
|
||||
flip_filter=False,
|
||||
gain=1,
|
||||
separable=None):
|
||||
r"""Convenience function to setup 2D FIR filter for `upfirdn2d()`.
|
||||
|
||||
Args:
|
||||
f: Torch tensor, numpy array, or python list of the shape
|
||||
`[filter_height, filter_width]` (non-separable),
|
||||
`[filter_taps]` (separable),
|
||||
`[]` (impulse), or
|
||||
`None` (identity).
|
||||
device: Result device (default: cpu).
|
||||
normalize: Normalize the filter so that it retains the magnitude
|
||||
for constant input signal (DC)? (default: True).
|
||||
flip_filter: Flip the filter? (default: False).
|
||||
gain: Overall scaling factor for signal magnitude (default: 1).
|
||||
separable: Return a separable filter? (default: select automatically).
|
||||
|
||||
Returns:
|
||||
Float32 tensor of the shape
|
||||
`[filter_height, filter_width]` (non-separable) or
|
||||
`[filter_taps]` (separable).
|
||||
"""
|
||||
# Validate.
|
||||
if f is None:
|
||||
f = 1
|
||||
f = torch.as_tensor(f, dtype=torch.float32)
|
||||
assert f.ndim in [0, 1, 2]
|
||||
assert f.numel() > 0
|
||||
if f.ndim == 0:
|
||||
f = f[np.newaxis]
|
||||
|
||||
# Separable?
|
||||
if separable is None:
|
||||
separable = (f.ndim == 1 and f.numel() >= 8)
|
||||
if f.ndim == 1 and not separable:
|
||||
f = f.ger(f)
|
||||
assert f.ndim == (1 if separable else 2)
|
||||
|
||||
# Apply normalize, flip, gain, and device.
|
||||
if normalize:
|
||||
f /= f.sum()
|
||||
if flip_filter:
|
||||
f = f.flip(list(range(f.ndim)))
|
||||
f = f * (gain**(f.ndim / 2))
|
||||
f = f.to(device=device)
|
||||
return f
|
||||
|
||||
|
||||
def upfirdn2d(x,
|
||||
f,
|
||||
up=1,
|
||||
down=1,
|
||||
padding=0,
|
||||
flip_filter=False,
|
||||
gain=1,
|
||||
impl='cuda'):
|
||||
r"""Pad, upsample, filter, and downsample a batch of 2D images.
|
||||
|
||||
Performs the following sequence of operations for each channel:
|
||||
|
||||
1. Upsample the image by inserting N-1 zeros after each pixel (`up`).
|
||||
|
||||
2. Pad the image with the specified number of zeros on each side (`padding`).
|
||||
Negative padding corresponds to cropping the image.
|
||||
|
||||
3. Convolve the image with the specified 2D FIR filter (`f`), shrinking it
|
||||
so that the footprint of all output pixels lies within the input image.
|
||||
|
||||
4. Downsample the image by keeping every Nth pixel (`down`).
|
||||
|
||||
This sequence of operations bears close resemblance to scipy.signal.upfirdn().
|
||||
The fused op is considerably more efficient than performing the same calculation
|
||||
using standard PyTorch ops. It supports gradients of arbitrary order.
|
||||
|
||||
Args:
|
||||
x: Float32/float64/float16 input tensor of the shape
|
||||
`[batch_size, num_channels, in_height, in_width]`.
|
||||
f: Float32 FIR filter of the shape
|
||||
`[filter_height, filter_width]` (non-separable),
|
||||
`[filter_taps]` (separable), or
|
||||
`None` (identity).
|
||||
up: Integer upsampling factor. Can be a single int or a list/tuple
|
||||
`[x, y]` (default: 1).
|
||||
down: Integer downsampling factor. Can be a single int or a list/tuple
|
||||
`[x, y]` (default: 1).
|
||||
padding: Padding with respect to the upsampled image. Can be a single number
|
||||
or a list/tuple `[x, y]` or `[x_before, x_after, y_before, y_after]`
|
||||
(default: 0).
|
||||
flip_filter: False = convolution, True = correlation (default: False).
|
||||
gain: Overall scaling factor for signal magnitude (default: 1).
|
||||
impl: Implementation to use. Can be `'ref'` or `'cuda'` (default: `'cuda'`).
|
||||
|
||||
Returns:
|
||||
Tensor of the shape `[batch_size, num_channels, out_height, out_width]`.
|
||||
"""
|
||||
assert isinstance(x, torch.Tensor)
|
||||
assert impl in ['ref', 'cuda']
|
||||
if impl == 'cuda' and x.device.type == 'cuda' and _init():
|
||||
return _upfirdn2d_cuda(
|
||||
up=up,
|
||||
down=down,
|
||||
padding=padding,
|
||||
flip_filter=flip_filter,
|
||||
gain=gain).apply(x, f)
|
||||
return _upfirdn2d_ref(
|
||||
x,
|
||||
f,
|
||||
up=up,
|
||||
down=down,
|
||||
padding=padding,
|
||||
flip_filter=flip_filter,
|
||||
gain=gain)
|
||||
|
||||
|
||||
@misc.profiled_function
|
||||
def _upfirdn2d_ref(x, f, up=1, down=1, padding=0, flip_filter=False, gain=1):
|
||||
"""Slow reference implementation of `upfirdn2d()` using standard PyTorch ops.
|
||||
"""
|
||||
# Validate arguments.
|
||||
assert isinstance(x, torch.Tensor) and x.ndim == 4
|
||||
if f is None:
|
||||
f = torch.ones([1, 1], dtype=torch.float32, device=x.device)
|
||||
assert isinstance(f, torch.Tensor) and f.ndim in [1, 2]
|
||||
assert f.dtype == torch.float32 and not f.requires_grad
|
||||
batch_size, num_channels, in_height, in_width = x.shape
|
||||
upx, upy = _parse_scaling(up)
|
||||
downx, downy = _parse_scaling(down)
|
||||
padx0, padx1, pady0, pady1 = _parse_padding(padding)
|
||||
|
||||
# Check that upsampled buffer is not smaller than the filter.
|
||||
upW = in_width * upx + padx0 + padx1
|
||||
upH = in_height * upy + pady0 + pady1
|
||||
assert upW >= f.shape[-1] and upH >= f.shape[0]
|
||||
|
||||
# Upsample by inserting zeros.
|
||||
x = x.reshape([batch_size, num_channels, in_height, 1, in_width, 1])
|
||||
x = torch.nn.functional.pad(x, [0, upx - 1, 0, 0, 0, upy - 1])
|
||||
x = x.reshape([batch_size, num_channels, in_height * upy, in_width * upx])
|
||||
|
||||
# Pad or crop.
|
||||
x = torch.nn.functional.pad(
|
||||
x, [max(padx0, 0),
|
||||
max(padx1, 0),
|
||||
max(pady0, 0),
|
||||
max(pady1, 0)])
|
||||
x = x[:, :,
|
||||
max(-pady0, 0):x.shape[2] - max(-pady1, 0),
|
||||
max(-padx0, 0):x.shape[3] - max(-padx1, 0)]
|
||||
|
||||
# Setup filter.
|
||||
f = f * (gain**(f.ndim / 2))
|
||||
f = f.to(x.dtype)
|
||||
if not flip_filter:
|
||||
f = f.flip(list(range(f.ndim)))
|
||||
|
||||
# Convolve with the filter.
|
||||
f = f[np.newaxis, np.newaxis].repeat([num_channels, 1] + [1] * f.ndim)
|
||||
if f.ndim == 4:
|
||||
x = conv2d_gradfix.conv2d(input=x, weight=f, groups=num_channels)
|
||||
else:
|
||||
x = conv2d_gradfix.conv2d(
|
||||
input=x, weight=f.unsqueeze(2), groups=num_channels)
|
||||
x = conv2d_gradfix.conv2d(
|
||||
input=x, weight=f.unsqueeze(3), groups=num_channels)
|
||||
|
||||
# Downsample by throwing away pixels.
|
||||
x = x[:, :, ::downy, ::downx]
|
||||
return x
|
||||
|
||||
|
||||
_upfirdn2d_cuda_cache = dict()
|
||||
|
||||
|
||||
def _upfirdn2d_cuda(up=1, down=1, padding=0, flip_filter=False, gain=1):
|
||||
"""Fast CUDA implementation of `upfirdn2d()` using custom ops.
|
||||
"""
|
||||
# Parse arguments.
|
||||
upx, upy = _parse_scaling(up)
|
||||
downx, downy = _parse_scaling(down)
|
||||
padx0, padx1, pady0, pady1 = _parse_padding(padding)
|
||||
|
||||
# Lookup from cache.
|
||||
key = (upx, upy, downx, downy, padx0, padx1, pady0, pady1, flip_filter,
|
||||
gain)
|
||||
if key in _upfirdn2d_cuda_cache:
|
||||
return _upfirdn2d_cuda_cache[key]
|
||||
|
||||
# Forward op.
|
||||
class Upfirdn2dCuda(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, x, f): # pylint: disable=arguments-differ
|
||||
assert isinstance(x, torch.Tensor) and x.ndim == 4
|
||||
if f is None:
|
||||
f = torch.ones([1, 1], dtype=torch.float32, device=x.device)
|
||||
if f.ndim == 1 and f.shape[0] == 1:
|
||||
f = f.square().unsqueeze(
|
||||
0) # Convert separable-1 into full-1x1.
|
||||
assert isinstance(f, torch.Tensor) and f.ndim in [1, 2]
|
||||
y = x
|
||||
if f.ndim == 2:
|
||||
y = _plugin.upfirdn2d(y, f, upx, upy, downx, downy, padx0,
|
||||
padx1, pady0, pady1, flip_filter, gain)
|
||||
else:
|
||||
y = _plugin.upfirdn2d(y, f.unsqueeze(0), upx, 1, downx, 1,
|
||||
padx0, padx1, 0, 0, flip_filter, 1.0)
|
||||
y = _plugin.upfirdn2d(y, f.unsqueeze(1), 1, upy, 1, downy, 0,
|
||||
0, pady0, pady1, flip_filter, gain)
|
||||
ctx.save_for_backward(f)
|
||||
ctx.x_shape = x.shape
|
||||
return y
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, dy): # pylint: disable=arguments-differ
|
||||
f, = ctx.saved_tensors
|
||||
_, _, ih, iw = ctx.x_shape
|
||||
_, _, oh, ow = dy.shape
|
||||
fw, fh = _get_filter_size(f)
|
||||
p = [
|
||||
fw - padx0 - 1,
|
||||
iw * upx - ow * downx + padx0 - upx + 1,
|
||||
fh - pady0 - 1,
|
||||
ih * upy - oh * downy + pady0 - upy + 1,
|
||||
]
|
||||
dx = None
|
||||
df = None
|
||||
|
||||
if ctx.needs_input_grad[0]:
|
||||
dx = _upfirdn2d_cuda(
|
||||
up=down,
|
||||
down=up,
|
||||
padding=p,
|
||||
flip_filter=(not flip_filter),
|
||||
gain=gain).apply(dy, f)
|
||||
|
||||
assert not ctx.needs_input_grad[1]
|
||||
return dx, df
|
||||
|
||||
# Add to cache.
|
||||
_upfirdn2d_cuda_cache[key] = Upfirdn2dCuda
|
||||
return Upfirdn2dCuda
|
||||
|
||||
|
||||
def filter2d(x, f, padding=0, flip_filter=False, gain=1, impl='cuda'):
|
||||
r"""Filter a batch of 2D images using the given 2D FIR filter.
|
||||
|
||||
By default, the result is padded so that its shape matches the input.
|
||||
User-specified padding is applied on top of that, with negative values
|
||||
indicating cropping. Pixels outside the image are assumed to be zero.
|
||||
|
||||
Args:
|
||||
x: Float32/float64/float16 input tensor of the shape
|
||||
`[batch_size, num_channels, in_height, in_width]`.
|
||||
f: Float32 FIR filter of the shape
|
||||
`[filter_height, filter_width]` (non-separable),
|
||||
`[filter_taps]` (separable), or
|
||||
`None` (identity).
|
||||
padding: Padding with respect to the output. Can be a single number or a
|
||||
list/tuple `[x, y]` or `[x_before, x_after, y_before, y_after]`
|
||||
(default: 0).
|
||||
flip_filter: False = convolution, True = correlation (default: False).
|
||||
gain: Overall scaling factor for signal magnitude (default: 1).
|
||||
impl: Implementation to use. Can be `'ref'` or `'cuda'` (default: `'cuda'`).
|
||||
|
||||
Returns:
|
||||
Tensor of the shape `[batch_size, num_channels, out_height, out_width]`.
|
||||
"""
|
||||
padx0, padx1, pady0, pady1 = _parse_padding(padding)
|
||||
fw, fh = _get_filter_size(f)
|
||||
p = [
|
||||
padx0 + fw // 2,
|
||||
padx1 + (fw - 1) // 2,
|
||||
pady0 + fh // 2,
|
||||
pady1 + (fh - 1) // 2,
|
||||
]
|
||||
return upfirdn2d(
|
||||
x, f, padding=p, flip_filter=flip_filter, gain=gain, impl=impl)
|
||||
|
||||
|
||||
def upsample2d(x, f, up=2, padding=0, flip_filter=False, gain=1, impl='cuda'):
|
||||
r"""Upsample a batch of 2D images using the given 2D FIR filter.
|
||||
|
||||
By default, the result is padded so that its shape is a multiple of the input.
|
||||
User-specified padding is applied on top of that, with negative values
|
||||
indicating cropping. Pixels outside the image are assumed to be zero.
|
||||
|
||||
Args:
|
||||
x: Float32/float64/float16 input tensor of the shape
|
||||
`[batch_size, num_channels, in_height, in_width]`.
|
||||
f: Float32 FIR filter of the shape
|
||||
`[filter_height, filter_width]` (non-separable),
|
||||
`[filter_taps]` (separable), or
|
||||
`None` (identity).
|
||||
up: Integer upsampling factor. Can be a single int or a list/tuple
|
||||
`[x, y]` (default: 1).
|
||||
padding: Padding with respect to the output. Can be a single number or a
|
||||
list/tuple `[x, y]` or `[x_before, x_after, y_before, y_after]`
|
||||
(default: 0).
|
||||
flip_filter: False = convolution, True = correlation (default: False).
|
||||
gain: Overall scaling factor for signal magnitude (default: 1).
|
||||
impl: Implementation to use. Can be `'ref'` or `'cuda'` (default: `'cuda'`).
|
||||
|
||||
Returns:
|
||||
Tensor of the shape `[batch_size, num_channels, out_height, out_width]`.
|
||||
"""
|
||||
upx, upy = _parse_scaling(up)
|
||||
padx0, padx1, pady0, pady1 = _parse_padding(padding)
|
||||
fw, fh = _get_filter_size(f)
|
||||
p = [
|
||||
padx0 + (fw + upx - 1) // 2,
|
||||
padx1 + (fw - upx) // 2,
|
||||
pady0 + (fh + upy - 1) // 2,
|
||||
pady1 + (fh - upy) // 2,
|
||||
]
|
||||
return upfirdn2d(
|
||||
x,
|
||||
f,
|
||||
up=up,
|
||||
padding=p,
|
||||
flip_filter=flip_filter,
|
||||
gain=gain * upx * upy,
|
||||
impl=impl)
|
||||
|
||||
|
||||
def downsample2d(x,
|
||||
f,
|
||||
down=2,
|
||||
padding=0,
|
||||
flip_filter=False,
|
||||
gain=1,
|
||||
impl='cuda'):
|
||||
r"""Downsample a batch of 2D images using the given 2D FIR filter.
|
||||
|
||||
By default, the result is padded so that its shape is a fraction of the input.
|
||||
User-specified padding is applied on top of that, with negative values
|
||||
indicating cropping. Pixels outside the image are assumed to be zero.
|
||||
|
||||
Args:
|
||||
x: Float32/float64/float16 input tensor of the shape
|
||||
`[batch_size, num_channels, in_height, in_width]`.
|
||||
f: Float32 FIR filter of the shape
|
||||
`[filter_height, filter_width]` (non-separable),
|
||||
`[filter_taps]` (separable), or
|
||||
`None` (identity).
|
||||
down: Integer downsampling factor. Can be a single int or a list/tuple
|
||||
`[x, y]` (default: 1).
|
||||
padding: Padding with respect to the input. Can be a single number or a
|
||||
list/tuple `[x, y]` or `[x_before, x_after, y_before, y_after]`
|
||||
(default: 0).
|
||||
flip_filter: False = convolution, True = correlation (default: False).
|
||||
gain: Overall scaling factor for signal magnitude (default: 1).
|
||||
impl: Implementation to use. Can be `'ref'` or `'cuda'` (default: `'cuda'`).
|
||||
|
||||
Returns:
|
||||
Tensor of the shape `[batch_size, num_channels, out_height, out_width]`.
|
||||
"""
|
||||
downx, downy = _parse_scaling(down)
|
||||
padx0, padx1, pady0, pady1 = _parse_padding(padding)
|
||||
fw, fh = _get_filter_size(f)
|
||||
p = [
|
||||
padx0 + (fw - downx + 1) // 2,
|
||||
padx1 + (fw - downx) // 2,
|
||||
pady0 + (fh - downy + 1) // 2,
|
||||
pady1 + (fh - downy) // 2,
|
||||
]
|
||||
return upfirdn2d(
|
||||
x,
|
||||
f,
|
||||
down=down,
|
||||
padding=p,
|
||||
flip_filter=flip_filter,
|
||||
gain=gain,
|
||||
impl=impl)
|
||||
@@ -0,0 +1,253 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2021-2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
|
||||
#
|
||||
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
|
||||
# property and proprietary rights in and to this material, related
|
||||
# documentation and any modifications thereto. Any use, reproduction,
|
||||
# disclosure or distribution of this material and related documentation
|
||||
# without an express license agreement from NVIDIA CORPORATION or
|
||||
# its affiliates is strictly prohibited.
|
||||
"""Facilities for pickling Python code alongside other data.
|
||||
|
||||
The pickled code is automatically imported into a separate Python module
|
||||
during unpickling. This way, any previously exported pickles will remain
|
||||
usable even if the original code is no longer available, or if the current
|
||||
version of the code is not consistent with what was originally pickled."""
|
||||
|
||||
import copy
|
||||
import inspect
|
||||
import io
|
||||
import pickle
|
||||
import sys
|
||||
import types
|
||||
import uuid
|
||||
|
||||
from .. import dnnlib
|
||||
|
||||
_version = 6 # internal version number
|
||||
_decorators = set() # {decorator_class, ...}
|
||||
_import_hooks = [] # [hook_function, ...]
|
||||
_module_to_src_dict = dict() # {module: src, ...}
|
||||
_src_to_module_dict = dict() # {src: module, ...}
|
||||
|
||||
|
||||
def persistent_class(orig_class):
|
||||
r"""Class decorator that extends a given class to save its source code
|
||||
when pickled.
|
||||
|
||||
Example:
|
||||
|
||||
from torch_utils import persistence
|
||||
|
||||
@persistence.persistent_class
|
||||
class MyNetwork(torch.nn.Module):
|
||||
def __init__(self, num_inputs, num_outputs):
|
||||
super().__init__()
|
||||
self.fc = MyLayer(num_inputs, num_outputs)
|
||||
...
|
||||
|
||||
@persistence.persistent_class
|
||||
class MyLayer(torch.nn.Module):
|
||||
...
|
||||
|
||||
When pickled, any instance of `MyNetwork` and `MyLayer` will save its
|
||||
source code alongside other internal state (e.g., parameters, buffers,
|
||||
and submodules). This way, any previously exported pickle will remain
|
||||
usable even if the class definitions have been modified or are no
|
||||
longer available.
|
||||
|
||||
The decorator saves the source code of the entire Python module
|
||||
containing the decorated class. It does *not* save the source code of
|
||||
any imported modules. Thus, the imported modules must be available
|
||||
during unpickling, also including `torch_utils.persistence` itself.
|
||||
|
||||
It is ok to call functions defined in the same module from the
|
||||
decorated class. However, if the decorated class depends on other
|
||||
classes defined in the same module, they must be decorated as well.
|
||||
This is illustrated in the above example in the case of `MyLayer`.
|
||||
|
||||
It is also possible to employ the decorator just-in-time before
|
||||
calling the constructor. For example:
|
||||
|
||||
cls = MyLayer
|
||||
if want_to_make_it_persistent:
|
||||
cls = persistence.persistent_class(cls)
|
||||
layer = cls(num_inputs, num_outputs)
|
||||
|
||||
As an additional feature, the decorator also keeps track of the
|
||||
arguments that were used to construct each instance of the decorated
|
||||
class. The arguments can be queried via `obj.init_args` and
|
||||
`obj.init_kwargs`, and they are automatically pickled alongside other
|
||||
object state. A typical use case is to first unpickle a previous
|
||||
instance of a persistent class, and then upgrade it to use the latest
|
||||
version of the source code:
|
||||
|
||||
with open('old_pickle.pkl', 'rb') as f:
|
||||
old_net = pickle.load(f)
|
||||
new_net = MyNetwork(*old_obj.init_args, **old_obj.init_kwargs)
|
||||
misc.copy_params_and_buffers(old_net, new_net, require_all=True)
|
||||
"""
|
||||
assert isinstance(orig_class, type)
|
||||
if is_persistent(orig_class):
|
||||
return orig_class
|
||||
|
||||
assert orig_class.__module__ in sys.modules
|
||||
orig_module = sys.modules[orig_class.__module__]
|
||||
orig_module_src = _module_to_src(orig_module)
|
||||
|
||||
class Decorator(orig_class):
|
||||
_orig_module_src = orig_module_src
|
||||
_orig_class_name = orig_class.__name__
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self._init_args = copy.deepcopy(args)
|
||||
self._init_kwargs = copy.deepcopy(kwargs)
|
||||
assert orig_class.__name__ in orig_module.__dict__
|
||||
_check_pickleable(self.__reduce__())
|
||||
|
||||
@property
|
||||
def init_args(self):
|
||||
return copy.deepcopy(self._init_args)
|
||||
|
||||
@property
|
||||
def init_kwargs(self):
|
||||
return dnnlib.EasyDict(copy.deepcopy(self._init_kwargs))
|
||||
|
||||
def __reduce__(self):
|
||||
fields = list(super().__reduce__())
|
||||
fields += [None] * max(3 - len(fields), 0)
|
||||
if fields[0] is not _reconstruct_persistent_obj:
|
||||
meta = dict(
|
||||
type='class',
|
||||
version=_version,
|
||||
module_src=self._orig_module_src,
|
||||
class_name=self._orig_class_name,
|
||||
state=fields[2])
|
||||
fields[0] = _reconstruct_persistent_obj # reconstruct func
|
||||
fields[1] = (meta, ) # reconstruct args
|
||||
fields[2] = None # state dict
|
||||
return tuple(fields)
|
||||
|
||||
Decorator.__name__ = orig_class.__name__
|
||||
_decorators.add(Decorator)
|
||||
return Decorator
|
||||
|
||||
|
||||
def is_persistent(obj):
|
||||
r"""Test whether the given object or class is persistent, i.e.,
|
||||
whether it will save its source code when pickled.
|
||||
"""
|
||||
try:
|
||||
if obj in _decorators:
|
||||
return True
|
||||
except TypeError:
|
||||
pass
|
||||
return type(obj) in _decorators # pylint: disable=unidiomatic-typecheck
|
||||
|
||||
|
||||
def import_hook(hook):
|
||||
r"""Register an import hook that is called whenever a persistent object
|
||||
is being unpickled. A typical use case is to patch the pickled source
|
||||
code to avoid errors and inconsistencies when the API of some imported
|
||||
module has changed.
|
||||
|
||||
The hook should have the following signature:
|
||||
|
||||
hook(meta) -> modified meta
|
||||
|
||||
`meta` is an instance of `dnnlib.EasyDict` with the following fields:
|
||||
|
||||
type: Type of the persistent object, e.g. `'class'`.
|
||||
version: Internal version number of `torch_utils.persistence`.
|
||||
module_src Original source code of the Python module.
|
||||
class_name: Class name in the original Python module.
|
||||
state: Internal state of the object.
|
||||
|
||||
Example:
|
||||
|
||||
@persistence.import_hook
|
||||
def wreck_my_network(meta):
|
||||
if meta.class_name == 'MyNetwork':
|
||||
print('MyNetwork is being imported. I will wreck it!')
|
||||
meta.module_src = meta.module_src.replace("True", "False")
|
||||
return meta
|
||||
"""
|
||||
assert callable(hook)
|
||||
_import_hooks.append(hook)
|
||||
|
||||
|
||||
def _reconstruct_persistent_obj(meta):
|
||||
r"""Hook that is called internally by the `pickle` module to unpickle
|
||||
a persistent object.
|
||||
"""
|
||||
meta = dnnlib.EasyDict(meta)
|
||||
meta.state = dnnlib.EasyDict(meta.state)
|
||||
for hook in _import_hooks:
|
||||
meta = hook(meta)
|
||||
assert meta is not None
|
||||
|
||||
assert meta.version == _version
|
||||
module = _src_to_module(meta.module_src)
|
||||
|
||||
assert meta.type == 'class'
|
||||
orig_class = module.__dict__[meta.class_name]
|
||||
decorator_class = persistent_class(orig_class)
|
||||
obj = decorator_class.__new__(decorator_class)
|
||||
|
||||
setstate = getattr(obj, '__setstate__', None)
|
||||
if callable(setstate):
|
||||
setstate(meta.state) # pylint: disable=not-callable
|
||||
else:
|
||||
obj.__dict__.update(meta.state)
|
||||
return obj
|
||||
|
||||
|
||||
def _module_to_src(module):
|
||||
r"""Query the source code of a given Python module.
|
||||
"""
|
||||
src = _module_to_src_dict.get(module, None)
|
||||
if src is None:
|
||||
src = inspect.getsource(module)
|
||||
_module_to_src_dict[module] = src
|
||||
_src_to_module_dict[src] = module
|
||||
return src
|
||||
|
||||
|
||||
def _src_to_module(src):
|
||||
r"""Get or create a Python module for the given source code.
|
||||
"""
|
||||
module = _src_to_module_dict.get(src, None)
|
||||
if module is None:
|
||||
module_name = '_imported_module_' + uuid.uuid4().hex
|
||||
module = types.ModuleType(module_name)
|
||||
sys.modules[module_name] = module
|
||||
_module_to_src_dict[module] = src
|
||||
_src_to_module_dict[src] = module
|
||||
exec(src, module.__dict__) # pylint: disable=exec-used
|
||||
return module
|
||||
|
||||
|
||||
def _check_pickleable(obj):
|
||||
r"""Check that the given object is pickleable, raising an exception if
|
||||
it is not. This function is expected to be considerably more efficient
|
||||
than actually pickling the object.
|
||||
"""
|
||||
|
||||
def recurse(obj):
|
||||
if isinstance(obj, (list, tuple, set)):
|
||||
return [recurse(x) for x in obj]
|
||||
if isinstance(obj, dict):
|
||||
return [[recurse(x), recurse(y)] for x, y in obj.items()]
|
||||
if isinstance(obj, (str, int, float, bool, bytes, bytearray)):
|
||||
return None # Python primitive types are pickleable.
|
||||
if f'{type(obj).__module__}.{type(obj).__name__}' in [
|
||||
'numpy.ndarray', 'torch.Tensor', 'torch.nn.parameter.Parameter'
|
||||
]:
|
||||
return None # NumPy arrays and PyTorch tensors are pickleable.
|
||||
if is_persistent(obj):
|
||||
return None # Persistent objects are pickleable, by virtue of the constructor check.
|
||||
return obj
|
||||
|
||||
with io.BytesIO() as f:
|
||||
pickle.dump(recurse(obj), f)
|
||||
@@ -758,6 +758,7 @@ TASK_OUTPUTS = {
|
||||
Tasks.nerf_recon_vq_compression: [OutputKeys.OUTPUT],
|
||||
Tasks.surface_recon_common: [OutputKeys.OUTPUT],
|
||||
Tasks.video_colorization: [OutputKeys.OUTPUT_VIDEO],
|
||||
Tasks.image_control_3d_portrait: [OutputKeys.OUTPUT],
|
||||
|
||||
# image quality assessment degradation result for single image
|
||||
# {
|
||||
|
||||
@@ -309,6 +309,10 @@ TASK_INPUTS = {
|
||||
InputKeys.IMAGE: InputType.IMAGE,
|
||||
'target_view': InputType.LIST
|
||||
},
|
||||
Tasks.image_control_3d_portrait: {
|
||||
InputKeys.IMAGE: InputType.IMAGE,
|
||||
'save_dir': InputType.TEXT
|
||||
},
|
||||
|
||||
# ============ nlp tasks ===================
|
||||
Tasks.chat: [
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from typing import Any, Dict
|
||||
|
||||
import numpy as np
|
||||
|
||||
from modelscope.metainfo import Pipelines
|
||||
from modelscope.outputs import OutputKeys
|
||||
from modelscope.pipelines.base import Input, Pipeline
|
||||
from modelscope.pipelines.builder import PIPELINES
|
||||
from modelscope.preprocessors import LoadImage
|
||||
from modelscope.utils.constant import ModelFile, Tasks
|
||||
from modelscope.utils.logger import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
@PIPELINES.register_module(
|
||||
Tasks.image_control_3d_portrait,
|
||||
module_name=Pipelines.image_control_3d_portrait)
|
||||
class ImageControl3dPortraitPipeline(Pipeline):
|
||||
""" Image control 3d portrait synthesis pipeline
|
||||
Example:
|
||||
|
||||
```python
|
||||
>>> from modelscope.pipelines import pipeline
|
||||
>>> image_control_3d_portrait = pipeline(Tasks.image_control_3d_portrait,
|
||||
'damo/cv_vit_image-control-3d-portrait-synthesis')
|
||||
>>> image_control_3d_portrait({
|
||||
'image_path': 'input.jpg', # input image path (str)
|
||||
'save_dir': 'save_dir', # save dir path (str)
|
||||
})
|
||||
>>>
|
||||
```
|
||||
"""
|
||||
|
||||
def __init__(self, model: str, **kwargs):
|
||||
"""
|
||||
use `model` to create image_control_3D_portrait pipeline for prediction
|
||||
Args:
|
||||
model: model id on modelscope hub.
|
||||
"""
|
||||
super().__init__(model=model, **kwargs)
|
||||
logger.info('image control 3D portrait synthesis model init done')
|
||||
|
||||
def preprocess(self, inputs: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return inputs
|
||||
|
||||
def forward(self, input: Dict[str, Any]) -> Dict[str, Any]:
|
||||
image_path = input['image']
|
||||
save_dir = input['save_dir']
|
||||
self.model.inference(image_path, save_dir)
|
||||
return {OutputKeys.OUTPUT: 'Done'}
|
||||
|
||||
def postprocess(self, inputs: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return inputs
|
||||
@@ -164,6 +164,7 @@ class CVTasks(object):
|
||||
nerf_recon_4k = 'nerf-recon-4k'
|
||||
nerf_recon_vq_compression = 'nerf-recon-vq-compression'
|
||||
surface_recon_common = 'surface-recon-common'
|
||||
image_control_3d_portrait = 'image-control-3d-portrait'
|
||||
|
||||
# vision efficient tuning
|
||||
vision_efficient_tuning = 'vision-efficient-tuning'
|
||||
|
||||
54
tests/pipelines/test_image_control_3d_portrait.py
Normal file
54
tests/pipelines/test_image_control_3d_portrait.py
Normal file
@@ -0,0 +1,54 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
|
||||
from modelscope.hub.api import HubApi
|
||||
from modelscope.hub.snapshot_download import snapshot_download
|
||||
from modelscope.outputs import OutputKeys
|
||||
from modelscope.pipelines import pipeline
|
||||
from modelscope.utils.constant import DownloadMode, Tasks
|
||||
from modelscope.utils.test_utils import test_level
|
||||
|
||||
|
||||
class ImageControl3dPortraitTest(unittest.TestCase):
|
||||
|
||||
def setUp(self) -> None:
|
||||
self.model_id = 'damo/cv_vit_image-control-3d-portrait-synthesis'
|
||||
self.test_image = 'data/test/images/image_control_3d_portrait.jpg'
|
||||
self.save_dir = 'exp'
|
||||
os.makedirs(self.save_dir, exist_ok=True)
|
||||
|
||||
@unittest.skipUnless(test_level() >= 2, 'skip test in current test level')
|
||||
def test_run_by_direct_model_download(self):
|
||||
model_dir = snapshot_download(self.model_id, revision='v1.1')
|
||||
print('model dir is: {}'.format(model_dir))
|
||||
image_control_3d_portrait = pipeline(
|
||||
Tasks.image_control_3d_portrait,
|
||||
model=model_dir,
|
||||
)
|
||||
image_control_3d_portrait(
|
||||
dict(image=self.test_image, save_dir=self.save_dir))
|
||||
|
||||
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
|
||||
def test_run_modelhub(self):
|
||||
image_control_3d_portrait = pipeline(
|
||||
Tasks.image_control_3d_portrait,
|
||||
model=self.model_id,
|
||||
)
|
||||
|
||||
image_control_3d_portrait(
|
||||
dict(image=self.test_image, save_dir=self.save_dir))
|
||||
print('image_control_3d_portrait.test_run_modelhub done')
|
||||
|
||||
@unittest.skipUnless(test_level() >= 2, 'skip test in current test level')
|
||||
def test_run_modelhub_default_model(self):
|
||||
image_control_3d_portrait = pipeline(Tasks.image_control_3d_portrait)
|
||||
image_control_3d_portrait(
|
||||
dict(image=self.test_image, save_dir=self.save_dir))
|
||||
print('image_control_3d_portrait.test_run_modelhub_default_model done')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user