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
| 161 | |
| 162 | |
| 163 | class 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__() |