File size: 5,468 Bytes
18e5c91
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
import torch
import numpy as np
import random
# from pytorch3d.transforms.rotation_conversions import quaternion_to_matrix


def gen_batch_ray_parellel(intrinsic,c2w,W,H):
    batch_size = intrinsic.shape[0]
    
    fx, fy, cx, cy = intrinsic[:,0,0].unsqueeze(1).unsqueeze(2), intrinsic[:,1,1].unsqueeze(1).unsqueeze(2), intrinsic[:,0,2].unsqueeze(1).unsqueeze(2), intrinsic[:,1,2].unsqueeze(1).unsqueeze(2)
    i, j = torch.meshgrid(torch.linspace(0.5, W-0.5, W, device=c2w.device), torch.linspace(0.5, H-0.5, H, device=c2w.device))  # pytorch's meshgrid has indexing='ij'
    i = i.t()
    j = j.t()
    i = i.unsqueeze(0).repeat(batch_size,1,1)
    j = j.unsqueeze(0).repeat(batch_size,1,1)
    dirs = torch.stack([(i-cx)/fx, (j-cy)/fy, torch.ones_like(i)], -1)
    rays_d = torch.sum(dirs[..., np.newaxis, :] * c2w[:,np.newaxis,np.newaxis, :3,:3], -1)  # dot product, equals to: [c2w.dot(dir) for dir in dirs]
    rays_o = c2w[:, :3, -1].unsqueeze(1).unsqueeze(2).repeat(1,H,W,1)
    viewdir = rays_d/torch.norm(rays_d,dim=-1,keepdim=True)
    return rays_d, rays_o, viewdir


def intrinsic_transform(intrinsic, original_res, size, transform_mode):
    fx, fy, cx, cy = intrinsic[0,0], intrinsic[1,1], intrinsic[0,2], intrinsic[1,2]
    original_height = original_res[0]
    original_width = original_res[1]
    if transform_mode == 'resize':
        resize_height = size[0]
        resize_width = size[1]

        scale_height = resize_height / original_height
        scale_width = resize_width / original_width

        fx_new = fx * scale_width
        fy_new = fy * scale_height
        cx_new = cx * scale_width
        cy_new = cy * scale_height

    elif transform_mode == 'center_crop_resize':
        if original_height <= original_width:
            scale_ratio = min(size) / original_height
        else:
            scale_ratio = min(size) / original_width
        resize_height = scale_ratio * original_height
        resize_width = scale_ratio * original_width

        fx_new = fx * scale_ratio
        fy_new = fy * scale_ratio
        cx_new = cx * scale_ratio
        cy_new = cy * scale_ratio

        crop_height = size[0]
        crop_width = size[1]
        cx_new = cx_new * (crop_width / resize_width)
        cy_new = cy_new * (crop_height / resize_height)
    
    else:
        raise NotImplementedError('No such transformation mode for image!')
    
    return torch.tensor([[fx_new, 0, cx_new],
                         [0, fy_new, cy_new],
                         [0, 0, 1]])




def intrinsic_transform_batch(intrinsic, original_res, size, transform_mode):
    b = intrinsic.shape[0]
    fx, fy, cx, cy = intrinsic[:,0,0], intrinsic[:,1,1], intrinsic[:,0,2], intrinsic[:,1,2]
    original_height = original_res[0]
    original_width = original_res[1]
    if transform_mode == 'resize':
        resize_height = size[0]
        resize_width = size[1]

        scale_height = resize_height / original_height
        scale_width = resize_width / original_width

        fx_new = fx * scale_width
        fy_new = fy * scale_height
        cx_new = cx * scale_width
        cy_new = cy * scale_height

    elif transform_mode == 'center_crop_resize':
        if original_height <= original_width:
            scale_ratio = min(size) / original_height
        else:
            scale_ratio = min(size) / original_width
        resize_height = scale_ratio * original_height
        resize_width = scale_ratio * original_width

        fx_new = fx * scale_ratio
        fy_new = fy * scale_ratio
        cx_new = cx * scale_ratio
        cy_new = cy * scale_ratio

        crop_height = size[0]
        crop_width = size[1]
        cx_new = cx_new * (crop_width / resize_width)
        cy_new = cy_new * (crop_height / resize_height)
    
    else:
        raise NotImplementedError('No such transformation mode for image!')
    

    fx_expanded = fx_new
    fy_expanded = fy_new
    cx_expanded = cx_new
    cy_expanded = cy_new

    intrinsic_matrices = torch.zeros((b, 3, 3), dtype=fx.dtype, device=fx.device)
    intrinsic_matrices[:, 0, 0] = fx_expanded
    intrinsic_matrices[:, 1, 1] = fy_expanded
    intrinsic_matrices[:, 0, 2] = cx_expanded
    intrinsic_matrices[:, 1, 2] = cy_expanded
    intrinsic_matrices[:, 2, 2] = 1
    
    return intrinsic_matrices


def gen_crop_config(tensor):
    h, w = tensor.shape[-2], tensor.shape[-1]
    h_start = random.randint(0,h//8)
    w_start = random.randint(0,w//8)
    h_crop = random.randint(7*h//8,h-h_start)
    w_crop = random.randint(7*w//8,w-w_start)
    return h_start, w_start, h_crop, w_crop


def crop_tensor(tensor, h_start, w_start, h_crop, w_crop):
    cropped_tensor = tensor[:,:,h_start:h_start+h_crop,w_start:w_start+w_crop]
    return cropped_tensor


def intrin_crop_transform(intrinsic, h_start, w_start):
    fx, fy, cx, cy = intrinsic[0][0], intrinsic[1][1], intrinsic[0][2], intrinsic[1][2]
    cx_new = cx - w_start
    cy_new = cy - h_start
    return torch.tensor([[fx,0,cx_new],[0,fy,cy_new],[0,0,1]])

# def get_transformation_matrix_from_quat(xyz_quat):
#     ### xyz_quat: tensor, (b, 7)
#     rot_quat = xyz_quat[:, 3:]
#     ### in pytorch3d, quaternion_to_matrix takes wxyz-quat as input
#     rot_quat = rot_quat[:, [3,0,1,2]]
#     rot = quaternion_to_matrix(rot_quat)
#     trans = xyz_quat[:, :3]
#     output = torch.eye(4).unsqueeze(0).repeat(xyz_quat.shape[0], 1, 1)
#     output[:,:3,:3] = rot
#     output[:,:3, 3] = trans
#     return output