(self,
D=8,
W=256,
d_in=3,
d_in_view=3,
multires=0,
multires_view=0,
output_ch=4,
skips=[4],
use_viewdirs=False)
| 175 | # This implementation is borrowed from nerf-pytorch: https://github.com/yenchenlin/nerf-pytorch |
| 176 | class NeRF(nn.Module): |
| 177 | def __init__(self, |
| 178 | D=8, |
| 179 | W=256, |
| 180 | d_in=3, |
| 181 | d_in_view=3, |
| 182 | multires=0, |
| 183 | multires_view=0, |
| 184 | output_ch=4, |
| 185 | skips=[4], |
| 186 | use_viewdirs=False): |
| 187 | super(NeRF, self).__init__() |
| 188 | self.D = D |
| 189 | self.W = W |
| 190 | self.d_in = d_in |
| 191 | self.d_in_view = d_in_view |
| 192 | self.input_ch = 3 |
| 193 | self.input_ch_view = 3 |
| 194 | self.embed_fn = None |
| 195 | self.embed_fn_view = None |
| 196 | |
| 197 | if multires > 0: |
| 198 | embed_fn, input_ch = get_embedder(multires, input_dims=d_in) |
| 199 | self.embed_fn = embed_fn |
| 200 | self.input_ch = input_ch |
| 201 | |
| 202 | if multires_view > 0: |
| 203 | embed_fn_view, input_ch_view = get_embedder(multires_view, input_dims=d_in_view) |
| 204 | self.embed_fn_view = embed_fn_view |
| 205 | self.input_ch_view = input_ch_view |
| 206 | |
| 207 | self.skips = skips |
| 208 | self.use_viewdirs = use_viewdirs |
| 209 | |
| 210 | self.pts_linears = nn.ModuleList( |
| 211 | [nn.Linear(self.input_ch, W)] + |
| 212 | [nn.Linear(W, W) if i not in self.skips else nn.Linear(W + self.input_ch, W) for i in range(D - 1)]) |
| 213 | |
| 214 | ### Implementation according to the official code release |
| 215 | ### (https://github.com/bmild/nerf/blob/master/run_nerf_helpers.py#L104-L105) |
| 216 | self.views_linears = nn.ModuleList([nn.Linear(self.input_ch_view + W, W // 2)]) |
| 217 | |
| 218 | ### Implementation according to the paper |
| 219 | # self.views_linears = nn.ModuleList( |
| 220 | # [nn.Linear(input_ch_views + W, W//2)] + [nn.Linear(W//2, W//2) for i in range(D//2)]) |
| 221 | |
| 222 | if use_viewdirs: |
| 223 | self.feature_linear = nn.Linear(W, W) |
| 224 | self.alpha_linear = nn.Linear(W, 1) |
| 225 | self.rgb_linear = nn.Linear(W // 2, 3) |
| 226 | else: |
| 227 | self.output_linear = nn.Linear(W, output_ch) |
| 228 | |
| 229 | def forward(self, input_pts, input_views): |
| 230 | if self.embed_fn is not None: |
no test coverage detected