MCPcopy Create free account
hub / github.com/apple/ml-pointersect / NetworkVRNN

Class NetworkVRNN

cdslib/core/nn/modules/vrnn.py:163–574  ·  view source on GitHub ↗

Variational RNN The inputs to the network are: - past outputs xs = [x0, x1, ..., x_{T-1}] - content sequence cs = [c1, ..., cN] - feature sequence fs = [f1, ..., fM] The output of the network: - ps = [p1, p2, ..., pT] - zs = [z1, z2, ..., zT] sequence of sampled

Source from the content-addressed store, hash-verified

161
162
163class NetworkVRNN(nn.Module):
164 """
165 Variational RNN
166
167 The inputs to the network are:
168
169 - past outputs xs = [x0, x1, ..., x_{T-1}]
170 - content sequence cs = [c1, ..., cN]
171 - feature sequence fs = [f1, ..., fM]
172
173 The output of the network:
174
175 - ps = [p1, p2, ..., pT]
176 - zs = [z1, z2, ..., zT] sequence of sampled latent variable
177
178 The network is composed of four parts:
179
180 - A :py:class:`NetworkGraves` that handles all the rendering, or P(pt | zt, x0, ..., x_{t-1}, cs)
181 - A multi-head dot-product attention that uses current hidden states to focus and extract info from fs
182 - A fully connected net to model the posterior q(zt | fs, cs, x0, ..., x_{t-1})
183 - A fully connected net to model the prior p(zt | cs, x0, ..., x_{t-1})
184
185 It does not include (since they are application-dependent):
186
187 - loss
188 - sampling method
189
190 Methods:
191 1. forward
192
193 - inputs:
194
195 - [x0, x1, ..., x_{T-1}]
196 - [c1, c2, ..., cN]
197 - [f1, f2, ..., fM]
198 - initial hidden states
199
200 - outputs:
201
202 - [p1, p2, ..., p_T]
203 - [\hat{c1}, ... \hat{cT}]
204 - final hidden states
205
206 2. get_all_zero_hidden_states
207 3. compute_zs
208 """
209
210 def __init__(self, param_dict: ParamVRNN = None, **kwargs):
211 """
212 Create a Variational Recurrent Neural Network (VRNN) model.
213
214 Args:
215 param_dict:
216 A :py:class:`ParamVRNN` object to define the hyper-parameters of the network.
217 kwargs:
218 If param_dict is None, you can directly provide keyword arguments of :py:class:`ParamVRNN` here.
219 """
220 super().__init__()

Callers 1

test_1Method · 0.90

Calls

no outgoing calls

Tested by 1

test_1Method · 0.72