MCPcopy Create free account
hub / github.com/Totoro97/NeuS / __init__

Method __init__

models/fields.py:177–227  ·  view source on GitHub ↗
(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)

Source from the content-addressed store, hash-verified

175# This implementation is borrowed from nerf-pytorch: https://github.com/yenchenlin/nerf-pytorch
176class 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:

Callers 3

__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45

Calls 1

get_embedderFunction · 0.90

Tested by

no test coverage detected