(self, image_size, obj_filename, uv_size=256, rasterizer_type='standard')
| 160 | |
| 161 | class SRenderY(nn.Module): |
| 162 | def __init__(self, image_size, obj_filename, uv_size=256, rasterizer_type='standard'): |
| 163 | super(SRenderY, self).__init__() |
| 164 | self.image_size = image_size |
| 165 | self.uv_size = uv_size |
| 166 | self.rasterizer_type=rasterizer_type |
| 167 | |
| 168 | if rasterizer_type == 'pytorch3d': |
| 169 | self.rasterizer = Pytorch3dRasterizer(image_size) |
| 170 | self.uv_rasterizer = Pytorch3dRasterizer(uv_size) |
| 171 | verts, faces, aux = load_obj(obj_filename) |
| 172 | uvcoords = aux.verts_uvs[None, ...] # (N, V, 2) |
| 173 | uvfaces = faces.textures_idx[None, ...] # (N, F, 3) |
| 174 | faces = faces.verts_idx[None, ...] |
| 175 | elif rasterizer_type == 'standard': |
| 176 | self.rasterizer = StandardRasterizer(image_size) |
| 177 | self.uv_rasterizer = StandardRasterizer(uv_size) |
| 178 | verts, uvcoords, faces, uvfaces = load_obj(obj_filename) |
| 179 | verts = verts[None, ...] |
| 180 | uvcoords = uvcoords[None, ...] |
| 181 | faces = faces[None, ...] |
| 182 | uvfaces = uvfaces[None, ...] |
| 183 | else: |
| 184 | NotImplementedError |
| 185 | |
| 186 | # faces |
| 187 | dense_triangles = util.generate_triangles(uv_size, uv_size) |
| 188 | self.register_buffer('dense_faces', torch.from_numpy( |
| 189 | dense_triangles).long()[None, :, :]) |
| 190 | self.register_buffer('faces', faces) |
| 191 | self.register_buffer('raw_uvcoords', uvcoords) |
| 192 | |
| 193 | # uv coords |
| 194 | uvcoords = torch.cat( |
| 195 | [uvcoords, uvcoords[:, :, 0:1]*0.+1.], -1) # [bz, ntv, 3] |
| 196 | uvcoords = uvcoords*2 - 1 |
| 197 | uvcoords[..., 1] = -uvcoords[..., 1] |
| 198 | face_uvcoords = util.face_vertices(uvcoords, uvfaces) |
| 199 | self.register_buffer('uvcoords', uvcoords) |
| 200 | self.register_buffer('uvfaces', uvfaces) |
| 201 | self.register_buffer('face_uvcoords', face_uvcoords) |
| 202 | |
| 203 | # shape colors, for rendering shape overlay |
| 204 | colors = torch.tensor([180, 180, 180])[None, None, :].repeat( |
| 205 | 1, faces.max()+1, 1).float()/255. |
| 206 | face_colors = util.face_vertices(colors, faces) |
| 207 | self.register_buffer('vertex_colors', colors) |
| 208 | self.register_buffer('face_colors', face_colors) |
| 209 | |
| 210 | # SH factors for lighting |
| 211 | pi = np.pi |
| 212 | constant_factor = torch.tensor([1/np.sqrt(4*pi), ((2*pi)/3)*(np.sqrt(3/(4*pi))), ((2*pi)/3)*(np.sqrt(3/(4*pi))), |
| 213 | ((2*pi)/3)*(np.sqrt(3/(4*pi))), (pi/4)*(3) * |
| 214 | (np.sqrt(5/(12*pi))), (pi/4) * |
| 215 | (3)*(np.sqrt(5/(12*pi))), |
| 216 | (pi/4)*(3)*(np.sqrt(5/(12*pi))), (pi/4)*(3/2)*(np.sqrt(5/(12*pi))), (pi/4)*(1/2)*(np.sqrt(5/(4*pi)))]).float() |
| 217 | self.register_buffer('constant_factor', constant_factor) |
| 218 | |
| 219 | def forward(self, vertices, transformed_vertices, albedos, lights=None, light_type='point', background=None, h=None, w=None): |
no test coverage detected