# MULTIVI training fails before first epoch

**URL:** <https://discourse.scverse.org/t/multivi-training-fails-before-first-epoch/236>\
**Category:** scvi-tools\
**Tags:** multivi\
**Created:** [December 9, 2021, 7:52am UTC](https://discourse.scverse.org/t/multivi-training-fails-before-first-epoch/236 "2021-12-09T07:52:45Z")\
**Posts on this page:** 2\
**Page:** 1

<div class="post-metadata">

**Author:** ![MoritzTh](https://yyz1.discourse-cdn.com/flex035/user_avatar/discourse.scverse.org/moritzth/32/116_2.png) [@MoritzTh](https://discourse.scverse.org/u/MoritzTh)\
**Post date:** [December 9, 2021, 7:52am UTC](https://discourse.scverse.org/t/multivi-training-fails-before-first-epoch/236/1 "2021-12-09T07:52:45Z")

</div>

Hi, I am trying to integrate mostly paired scRNA and scATAC data following your MULTIVI tutorial. Creating the mvi anndata and setting up the model with

`scvi.model.MULTIVI.setup_anndata(adata_mvi, batch_key='modality')`

works fine. However, training the model with

```auto
mvi = scvi.model.MULTIVI(
    adata_mvi,
    n_genes=(adata_mvi.var['modality']=='Gene Expression').sum(),
    n_regions=(adata_mvi.var['modality']=='Peaks').sum(),
)
mvi.train()

```

results in the error:

```auto
GPU available: False, used: False
TPU available: False, using: 0 TPU cores
Epoch 1/500: 0%| | 0/500 [00:00<?, ?it/s]
---------------------------------------------------------------------------
ValueError Traceback (most recent call last)
<ipython-input-14-b4772eb1a555> in <module>
      4 n_regions=(adata_mvi.var['modality']=='Peaks').sum(),
      5 )
----> 6 mvi.train()

~/miniconda3/lib/python3.7/site-packages/scvi/model/_multivi.py in train(self, max_epochs, lr, use_gpu, train_size, validation_size, batch_size, weight_decay, eps, early_stopping, save_best, check_val_every_n_epoch, n_steps_kl_warmup, n_epochs_kl_warmup, adversarial_mixing, plan_kwargs, **kwargs)
    278 **kwargs,
    279 )
--> 280 return runner()
    281 
    282 @torch.no_grad()

~/miniconda3/lib/python3.7/site-packages/scvi/train/_trainrunner.py in __call__ (self)
     70 self.training_plan.n_obs_training = self.data_splitter.n_train
     71 
---> 72 self.trainer.fit(self.training_plan, self.data_splitter)
     73 self._update_history()
     74 

~/miniconda3/lib/python3.7/site-packages/scvi/train/_trainer.py in fit(self, *args, **kwargs)
    175 message="`LightningModule.configure_optimizers` returned `None`",
    176 )
--> 177 super().fit(*args, **kwargs)

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/trainer/trainer.py in fit(self, model, train_dataloader, val_dataloaders, datamodule)
    458 )
    459 
--> 460 self._run(model)
    461 
    462 assert self.state.stopped

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/trainer/trainer.py in _run(self, model)
    756 
    757 # dispatch `start_training` or `start_evaluating` or `start_predicting`
--> 758 self.dispatch()
    759 
    760 # plugin will finalized fitting (e.g. ddp_spawn will load trained model)

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/trainer/trainer.py in dispatch(self)
    797 self.accelerator.start_predicting(self)
    798 else:
--> 799 self.accelerator.start_training(self)
    800 
    801 def run_stage(self):

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/accelerators/accelerator.py in start_training(self, trainer)
     94 
     95 def start_training(self, trainer: 'pl.Trainer') -> None:
---> 96 self.training_type_plugin.start_training(trainer)
     97 
     98 def start_evaluating(self, trainer: 'pl.Trainer') -> None:

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/plugins/training_type/training_type_plugin.py in start_training(self, trainer)
    142 def start_training(self, trainer: 'pl.Trainer') -> None:
    143 # double dispatch to initiate the training loop
--> 144 self._results = trainer.run_stage()
    145 
    146 def start_evaluating(self, trainer: 'pl.Trainer') -> None:

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/trainer/trainer.py in run_stage(self)
    807 if self.predicting:
    808 return self.run_predict()
--> 809 return self.run_train()
    810 
    811 def _pre_training_routine(self):

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/trainer/trainer.py in run_train(self)
    869 with self.profiler.profile("run_training_epoch"):
    870 # run train epoch
--> 871 self.train_loop.run_training_epoch()
    872 
    873 if self.max_steps and self.max_steps <= self.global_step:

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/trainer/training_loop.py in run_training_epoch(self)
    497 # ------------------------------------
    498 with self.trainer.profiler.profile("run_training_batch"):
--> 499 batch_output = self.run_training_batch(batch, batch_idx, dataloader_idx)
    500 
    501 # when returning -1 from train_step, we end epoch early

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/trainer/training_loop.py in run_training_batch(self, batch, batch_idx, dataloader_idx)
    736 
    737 # optimizer step
--> 738 self.optimizer_step(optimizer, opt_idx, batch_idx, train_step_and_backward_closure)
    739 if len(self.trainer.optimizers) > 1:
    740 # revert back to previous state

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/trainer/training_loop.py in optimizer_step(self, optimizer, opt_idx, batch_idx, train_step_and_backward_closure)
    440 on_tpu=self.trainer._device_type == DeviceType.TPU and _TPU_AVAILABLE,
    441 using_native_amp=using_native_amp,
--> 442 using_lbfgs=is_lbfgs,
    443 )
    444 

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/core/lightning.py in optimizer_step(self, epoch, batch_idx, optimizer, optimizer_idx, optimizer_closure, on_tpu, using_native_amp, using_lbfgs)
   1401 
   1402 """
-> 1403 optimizer.step(closure=optimizer_closure)
   1404 
   1405 def optimizer_zero_grad(self, epoch: int, batch_idx: int, optimizer: Optimizer, optimizer_idx: int):

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/core/optimizer.py in step(self, closure, *args, **kwargs)
    212 profiler_name = f"optimizer_step_and_closure_{self._optimizer_idx}"
    213 
--> 214 self.__optimizer_step(*args, closure=closure, profiler_name=profiler_name, **kwargs)
    215 self._total_optimizer_step_calls += 1
    216 

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/core/optimizer.py in __optimizer_step(self, closure, profiler_name, **kwargs)
    132 
    133 with trainer.profiler.profile(profiler_name):
--> 134 trainer.accelerator.optimizer_step(optimizer, self._optimizer_idx, lambda_closure=closure, **kwargs)
    135 
    136 def step(self, *args, closure: Optional[Callable] = None, **kwargs):

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/accelerators/accelerator.py in optimizer_step(self, optimizer, opt_idx, lambda_closure, **kwargs)
    327 )
    328 if make_optimizer_step:
--> 329 self.run_optimizer_step(optimizer, opt_idx, lambda_closure, **kwargs)
    330 self.precision_plugin.post_optimizer_step(optimizer, opt_idx)
    331 self.training_type_plugin.post_optimizer_step(optimizer, opt_idx, **kwargs)

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/accelerators/accelerator.py in run_optimizer_step(self, optimizer, optimizer_idx, lambda_closure, **kwargs)
    334 self, optimizer: Optimizer, optimizer_idx: int, lambda_closure: Callable, **kwargs: Any
    335 ) -> None:
--> 336 self.training_type_plugin.optimizer_step(optimizer, lambda_closure=lambda_closure, **kwargs)
    337 
    338 def optimizer_zero_grad(self, current_epoch: int, batch_idx: int, optimizer: Optimizer, opt_idx: int) -> None:

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/plugins/training_type/training_type_plugin.py in optimizer_step(self, optimizer, lambda_closure, **kwargs)
    191 
    192 def optimizer_step(self, optimizer: torch.optim.Optimizer, lambda_closure: Callable, **kwargs):
--> 193 optimizer.step(closure=lambda_closure, **kwargs)
    194 
    195 @property

~/miniconda3/lib/python3.7/site-packages/torch/optim/optimizer.py in wrapper(*args, **kwargs)
     86 profile_name = "Optimizer.step#{}.step".format(obj. __class__. __name__ )
     87 with torch.autograd.profiler.record_function(profile_name):
---> 88 return func(*args, **kwargs)
     89 return wrapper
     90 

~/miniconda3/lib/python3.7/site-packages/torch/autograd/grad_mode.py in decorate_context(*args, **kwargs)
     26 def decorate_context(*args, **kwargs):
     27 with self. __class__ ():
---> 28 return func(*args, **kwargs)
     29 return cast(F, decorate_context)
     30 

~/miniconda3/lib/python3.7/site-packages/torch/optim/adam.py in step(self, closure)
     90 if closure is not None:
     91 with torch.enable_grad():
---> 92 loss = closure()
     93 
     94 for group in self.param_groups:

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/trainer/training_loop.py in train_step_and_backward_closure()
    731 def train_step_and_backward_closure():
    732 result = self.training_step_and_backward(
--> 733 split_batch, batch_idx, opt_idx, optimizer, self.trainer.hiddens
    734 )
    735 return None if result is None else result.loss

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/trainer/training_loop.py in training_step_and_backward(self, split_batch, batch_idx, opt_idx, optimizer, hiddens)
    821 with self.trainer.profiler.profile("training_step_and_backward"):
    822 # lightning module hook
--> 823 result = self.training_step(split_batch, batch_idx, opt_idx, hiddens)
    824 self._curr_step_result = result
    825 

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/trainer/training_loop.py in training_step(self, split_batch, batch_idx, opt_idx, hiddens)
    288 model_ref._results = Result()
    289 with self.trainer.profiler.profile("training_step"):
--> 290 training_step_output = self.trainer.accelerator.training_step(args)
    291 self.trainer.accelerator.post_training_step()
    292 

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/accelerators/accelerator.py in training_step(self, args)
    202 
    203 with self.precision_plugin.train_step_context(), self.training_type_plugin.train_step_context():
--> 204 return self.training_type_plugin.training_step(*args)
    205 
    206 def post_training_step(self) -> None:

~/miniconda3/lib/python3.7/site-packages/pytorch_lightning/plugins/training_type/training_type_plugin.py in training_step(self, *args, **kwargs)
    153 
    154 def training_step(self, *args, **kwargs):
--> 155 return self.lightning_module.training_step(*args, **kwargs)
    156 
    157 def post_training_step(self):

~/miniconda3/lib/python3.7/site-packages/scvi/train/_trainingplans.py in training_step(self, batch, batch_idx, optimizer_idx)
    362 loss_kwargs = dict(kl_weight=self.kl_weight)
    363 inference_outputs, _, scvi_loss = self.forward(
--> 364 batch, loss_kwargs=loss_kwargs
    365 )
    366 loss = scvi_loss.loss

~/miniconda3/lib/python3.7/site-packages/scvi/train/_trainingplans.py in forward(self, *args, **kwargs)
    145 def forward(self, *args, **kwargs):
    146 """Passthrough to `model.forward()`."""
--> 147 return self.module(*args, **kwargs)
    148 
    149 def training_step(self, batch, batch_idx, optimizer_idx=0):

~/miniconda3/lib/python3.7/site-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs)
   1100 if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks
   1101 or _global_forward_hooks or _global_forward_pre_hooks):
-> 1102 return forward_call(*input, **kwargs)
   1103 # Do not call functions when jit is used
   1104 full_backward_hooks, non_full_backward_hooks = [], []

~/miniconda3/lib/python3.7/site-packages/scvi/module/base/_decorators.py in auto_transfer_args(self, *args, **kwargs)
     30 # decorator only necessary after training
     31 if self.training:
---> 32 return fn(self, *args, **kwargs)
     33 
     34 device = list(set(p.device for p in self.parameters()))

~/miniconda3/lib/python3.7/site-packages/scvi/module/base/_base_module.py in forward(self, tensors, get_inference_input_kwargs, get_generative_input_kwargs, inference_kwargs, generative_kwargs, loss_kwargs, compute_loss)
    143 tensors, **get_inference_input_kwargs
    144 )
--> 145 inference_outputs = self.inference( **inference_inputs,** inference_kwargs)
    146 generative_inputs = self._get_generative_input(
    147 tensors, inference_outputs, **get_generative_input_kwargs

~/miniconda3/lib/python3.7/site-packages/scvi/module/base/_decorators.py in auto_transfer_args(self, *args, **kwargs)
     30 # decorator only necessary after training
     31 if self.training:
---> 32 return fn(self, *args, **kwargs)
     33 
     34 device = list(set(p.device for p in self.parameters()))

~/miniconda3/lib/python3.7/site-packages/scvi/module/_multivae.py in inference(self, x, batch_index, cont_covs, cat_covs, n_samples)
    293 # Z Encoders
    294 qzm_acc, qzv_acc, z_acc = self.z_encoder_accessibility(
--> 295 encoder_input_accessibility, batch_index, *categorical_input
    296 )
    297 qzm_expr, qzv_expr, z_expr = self.z_encoder_expression(

~/miniconda3/lib/python3.7/site-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs)
   1100 if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks
   1101 or _global_forward_hooks or _global_forward_pre_hooks):
-> 1102 return forward_call(*input, **kwargs)
   1103 # Do not call functions when jit is used
   1104 full_backward_hooks, non_full_backward_hooks = [], []

~/miniconda3/lib/python3.7/site-packages/scvi/nn/_base_components.py in forward(self, x, *cat_list)
    292 q_m = self.mean_encoder(q)
    293 q_v = self.var_activation(self.var_encoder(q)) + self.var_eps
--> 294 latent = self.z_transformation(reparameterize_gaussian(q_m, q_v))
    295 return q_m, q_v, latent
    296 

~/miniconda3/lib/python3.7/site-packages/scvi/nn/_base_components.py in reparameterize_gaussian(mu, var)
     11 
     12 def reparameterize_gaussian(mu, var):
---> 13 return Normal(mu, var.sqrt()).rsample()
     14 
     15 

~/miniconda3/lib/python3.7/site-packages/torch/distributions/normal.py in __init__ (self, loc, scale, validate_args)
     48 else:
     49 batch_shape = self.loc.size()
---> 50 super(Normal, self). __init__ (batch_shape, validate_args=validate_args)
     51 
     52 def expand(self, batch_shape, _instance=None):

~/miniconda3/lib/python3.7/site-packages/torch/distributions/distribution.py in __init__ (self, batch_shape, event_shape, validate_args)
     54 if not valid.all():
     55 raise ValueError(
---> 56 f"Expected parameter {param} "
     57 f"({type(value). __name__ } of shape {tuple(value.shape)}) "
     58 f"of distribution {repr(self)} "

ValueError: Expected parameter loc (Tensor of shape (128, 19)) of distribution Normal(loc: torch.Size([128, 19]), scale: torch.Size([128, 19])) to satisfy the constraint Real(), but found invalid values:
tensor([[0.1109, 0.5803, 0.3902, ..., -0.4835, -0.8638, 0.0870],
        [0.7019, 0.4671, 0.4204, ..., 0.0102, -0.8212, 0.0126],
        [0.1285, 1.1512, -0.1905, ..., -0.4889, -0.0262, 0.0351],
        ...,
        [-0.0526, 0.3792, 0.6689, ..., 0.2085, 0.0496, 0.4914],
        [0.5148, 0.4604, 0.4606, ..., -0.0603, -0.3616, 0.4082],
        [0.5613, 0.6148, 0.5383, ..., -0.1725, -1.2356, -0.1374]],
       grad_fn=<AddmmBackward0>)

```

Can you help me with this?

Thanks in advance!  
Moritz

---

<div class="post-metadata">

**Author:** ![adamgayoso](https://yyz1.discourse-cdn.com/flex035/user_avatar/discourse.scverse.org/adamgayoso/32/100_2.png) [@adamgayoso](https://discourse.scverse.org/u/adamgayoso)\
**Post date:** [December 16, 2021, 4:21pm UTC](https://discourse.scverse.org/t/multivi-training-fails-before-first-epoch/236/2 "2021-12-16T16:21:20Z")

</div>

Discussion here:

> <https://github.com/YosefLab/scvi-tools/issues/1291>
>
> Hi, I am trying to integrate mostly paired scRNA and scATAC data following your …MULTIVI tutorial. Creating the mvi anndata and setting up the model with
> \`\`\`
> scvi.model.MULTIVI.setup\_anndata(adata\_mvi, batch\_key='modality')
> \`\`\`
> works fine. However, training the model with:
> \`\`\`
> mvi = scvi.model.MULTIVI(
> adata\_mvi,
> n\_genes=(adata\_mvi.var\['modality'\]=='Gene Expression').sum(),
> n\_regions=(adata\_mvi.var\['modality'\]=='Peaks').sum(),
> )
> mvi.train()
> \`\`\`
> results in the error:
> 
> \`\`\`pytb
> GPU available: False, used: False
> TPU available: False, using: 0 TPU cores
> Epoch 1/500: 0%| | 0/500 \[00:00\<?, ?it/s\]
> \---------------------------------------------------------------------------
> ValueError Traceback (most recent call last)
> \<ipython-input-14-b4772eb1a555\> in \<module\>
> 4 n\_regions=(adata\_mvi.var\['modality'\]=='Peaks').sum(),
> 5 )
> \----\> 6 mvi.train()
> 
> ~/miniconda3/lib/python3.7/site-packages/scvi/model/\_multivi.py in train(self, max\_epochs, lr, use\_gpu, train\_size, validation\_size, batch\_size, weight\_decay, eps, early\_stopping, save\_best, check\_val\_every\_n\_epoch, n\_steps\_kl\_warmup, n\_epochs\_kl\_warmup, adversarial\_mixing, plan\_kwargs, \*\*kwargs)
> 278 \*\*kwargs,
> 279 )
> \--\> 280 return runner()
> 281 
> 282 @torch.no\_grad()
> 
> ~/miniconda3/lib/python3.7/site-packages/scvi/train/\_trainrunner.py in \_\_call\_\_(self)
> 70 self.training\_plan.n\_obs\_training = self.data\_splitter.n\_train
> 71 
> \---\> 72 self.trainer.fit(self.training\_plan, self.data\_splitter)
> 73 self.\_update\_history()
> 74 
> 
> ~/miniconda3/lib/python3.7/site-packages/scvi/train/\_trainer.py in fit(self, \*args, \*\*kwargs)
> 175 message="\`LightningModule.configure\_optimizers\` returned \`None\`",
> 176 )
> \--\> 177 super().fit(\*args, \*\*kwargs)
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/trainer/trainer.py in fit(self, model, train\_dataloader, val\_dataloaders, datamodule)
> 458 )
> 459 
> \--\> 460 self.\_run(model)
> 461 
> 462 assert self.state.stopped
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/trainer/trainer.py in \_run(self, model)
> 756 
> 757 # dispatch \`start\_training\` or \`start\_evaluating\` or \`start\_predicting\`
> \--\> 758 self.dispatch()
> 759 
> 760 # plugin will finalized fitting (e.g. ddp\_spawn will load trained model)
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/trainer/trainer.py in dispatch(self)
> 797 self.accelerator.start\_predicting(self)
> 798 else:
> \--\> 799 self.accelerator.start\_training(self)
> 800 
> 801 def run\_stage(self):
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/accelerators/accelerator.py in start\_training(self, trainer)
> 94 
> 95 def start\_training(self, trainer: 'pl.Trainer') -\> None:
> \---\> 96 self.training\_type\_plugin.start\_training(trainer)
> 97 
> 98 def start\_evaluating(self, trainer: 'pl.Trainer') -\> None:
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/plugins/training\_type/training\_type\_plugin.py in start\_training(self, trainer)
> 142 def start\_training(self, trainer: 'pl.Trainer') -\> None:
> 143 # double dispatch to initiate the training loop
> \--\> 144 self.\_results = trainer.run\_stage()
> 145 
> 146 def start\_evaluating(self, trainer: 'pl.Trainer') -\> None:
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/trainer/trainer.py in run\_stage(self)
> 807 if self.predicting:
> 808 return self.run\_predict()
> \--\> 809 return self.run\_train()
> 810 
> 811 def \_pre\_training\_routine(self):
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/trainer/trainer.py in run\_train(self)
> 869 with self.profiler.profile("run\_training\_epoch"):
> 870 # run train epoch
> \--\> 871 self.train\_loop.run\_training\_epoch()
> 872 
> 873 if self.max\_steps and self.max\_steps \<= self.global\_step:
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/trainer/training\_loop.py in run\_training\_epoch(self)
> 497 # ------------------------------------
> 498 with self.trainer.profiler.profile("run\_training\_batch"):
> \--\> 499 batch\_output = self.run\_training\_batch(batch, batch\_idx, dataloader\_idx)
> 500 
> 501 # when returning -1 from train\_step, we end epoch early
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/trainer/training\_loop.py in run\_training\_batch(self, batch, batch\_idx, dataloader\_idx)
> 736 
> 737 # optimizer step
> \--\> 738 self.optimizer\_step(optimizer, opt\_idx, batch\_idx, train\_step\_and\_backward\_closure)
> 739 if len(self.trainer.optimizers) \> 1:
> 740 # revert back to previous state
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/trainer/training\_loop.py in optimizer\_step(self, optimizer, opt\_idx, batch\_idx, train\_step\_and\_backward\_closure)
> 440 on\_tpu=self.trainer.\_device\_type == DeviceType.TPU and \_TPU\_AVAILABLE,
> 441 using\_native\_amp=using\_native\_amp,
> \--\> 442 using\_lbfgs=is\_lbfgs,
> 443 )
> 444 
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/core/lightning.py in optimizer\_step(self, epoch, batch\_idx, optimizer, optimizer\_idx, optimizer\_closure, on\_tpu, using\_native\_amp, using\_lbfgs)
> 1401 
> 1402 """
> \-\> 1403 optimizer.step(closure=optimizer\_closure)
> 1404 
> 1405 def optimizer\_zero\_grad(self, epoch: int, batch\_idx: int, optimizer: Optimizer, optimizer\_idx: int):
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/core/optimizer.py in step(self, closure, \*args, \*\*kwargs)
> 212 profiler\_name = f"optimizer\_step\_and\_closure\_{self.\_optimizer\_idx}"
> 213 
> \--\> 214 self.\_\_optimizer\_step(\*args, closure=closure, profiler\_name=profiler\_name, \*\*kwargs)
> 215 self.\_total\_optimizer\_step\_calls += 1
> 216 
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/core/optimizer.py in \_\_optimizer\_step(self, closure, profiler\_name, \*\*kwargs)
> 132 
> 133 with trainer.profiler.profile(profiler\_name):
> \--\> 134 trainer.accelerator.optimizer\_step(optimizer, self.\_optimizer\_idx, lambda\_closure=closure, \*\*kwargs)
> 135 
> 136 def step(self, \*args, closure: Optional\[Callable\] = None, \*\*kwargs):
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/accelerators/accelerator.py in optimizer\_step(self, optimizer, opt\_idx, lambda\_closure, \*\*kwargs)
> 327 )
> 328 if make\_optimizer\_step:
> \--\> 329 self.run\_optimizer\_step(optimizer, opt\_idx, lambda\_closure, \*\*kwargs)
> 330 self.precision\_plugin.post\_optimizer\_step(optimizer, opt\_idx)
> 331 self.training\_type\_plugin.post\_optimizer\_step(optimizer, opt\_idx, \*\*kwargs)
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/accelerators/accelerator.py in run\_optimizer\_step(self, optimizer, optimizer\_idx, lambda\_closure, \*\*kwargs)
> 334 self, optimizer: Optimizer, optimizer\_idx: int, lambda\_closure: Callable, \*\*kwargs: Any
> 335 ) -\> None:
> \--\> 336 self.training\_type\_plugin.optimizer\_step(optimizer, lambda\_closure=lambda\_closure, \*\*kwargs)
> 337 
> 338 def optimizer\_zero\_grad(self, current\_epoch: int, batch\_idx: int, optimizer: Optimizer, opt\_idx: int) -\> None:
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/plugins/training\_type/training\_type\_plugin.py in optimizer\_step(self, optimizer, lambda\_closure, \*\*kwargs)
> 191 
> 192 def optimizer\_step(self, optimizer: torch.optim.Optimizer, lambda\_closure: Callable, \*\*kwargs):
> \--\> 193 optimizer.step(closure=lambda\_closure, \*\*kwargs)
> 194 
> 195 @property
> 
> ~/miniconda3/lib/python3.7/site-packages/torch/optim/optimizer.py in wrapper(\*args, \*\*kwargs)
> 86 profile\_name = "Optimizer.step#{}.step".format(obj.\_\_class\_\_.\_\_name\_\_)
> 87 with torch.autograd.profiler.record\_function(profile\_name):
> \---\> 88 return func(\*args, \*\*kwargs)
> 89 return wrapper
> 90 
> 
> ~/miniconda3/lib/python3.7/site-packages/torch/autograd/grad\_mode.py in decorate\_context(\*args, \*\*kwargs)
> 26 def decorate\_context(\*args, \*\*kwargs):
> 27 with self.\_\_class\_\_():
> \---\> 28 return func(\*args, \*\*kwargs)
> 29 return cast(F, decorate\_context)
> 30 
> 
> ~/miniconda3/lib/python3.7/site-packages/torch/optim/adam.py in step(self, closure)
> 90 if closure is not None:
> 91 with torch.enable\_grad():
> \---\> 92 loss = closure()
> 93 
> 94 for group in self.param\_groups:
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/trainer/training\_loop.py in train\_step\_and\_backward\_closure()
> 731 def train\_step\_and\_backward\_closure():
> 732 result = self.training\_step\_and\_backward(
> \--\> 733 split\_batch, batch\_idx, opt\_idx, optimizer, self.trainer.hiddens
> 734 )
> 735 return None if result is None else result.loss
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/trainer/training\_loop.py in training\_step\_and\_backward(self, split\_batch, batch\_idx, opt\_idx, optimizer, hiddens)
> 821 with self.trainer.profiler.profile("training\_step\_and\_backward"):
> 822 # lightning module hook
> \--\> 823 result = self.training\_step(split\_batch, batch\_idx, opt\_idx, hiddens)
> 824 self.\_curr\_step\_result = result
> 825 
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/trainer/training\_loop.py in training\_step(self, split\_batch, batch\_idx, opt\_idx, hiddens)
> 288 model\_ref.\_results = Result()
> 289 with self.trainer.profiler.profile("training\_step"):
> \--\> 290 training\_step\_output = self.trainer.accelerator.training\_step(args)
> 291 self.trainer.accelerator.post\_training\_step()
> 292 
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/accelerators/accelerator.py in training\_step(self, args)
> 202 
> 203 with self.precision\_plugin.train\_step\_context(), self.training\_type\_plugin.train\_step\_context():
> \--\> 204 return self.training\_type\_plugin.training\_step(\*args)
> 205 
> 206 def post\_training\_step(self) -\> None:
> 
> ~/miniconda3/lib/python3.7/site-packages/pytorch\_lightning/plugins/training\_type/training\_type\_plugin.py in training\_step(self, \*args, \*\*kwargs)
> 153 
> 154 def training\_step(self, \*args, \*\*kwargs):
> \--\> 155 return self.lightning\_module.training\_step(\*args, \*\*kwargs)
> 156 
> 157 def post\_training\_step(self):
> 
> ~/miniconda3/lib/python3.7/site-packages/scvi/train/\_trainingplans.py in training\_step(self, batch, batch\_idx, optimizer\_idx)
> 362 loss\_kwargs = dict(kl\_weight=self.kl\_weight)
> 363 inference\_outputs, \_, scvi\_loss = self.forward(
> \--\> 364 batch, loss\_kwargs=loss\_kwargs
> 365 )
> 366 loss = scvi\_loss.loss
> 
> ~/miniconda3/lib/python3.7/site-packages/scvi/train/\_trainingplans.py in forward(self, \*args, \*\*kwargs)
> 145 def forward(self, \*args, \*\*kwargs):
> 146 """Passthrough to \`model.forward()\`."""
> \--\> 147 return self.module(\*args, \*\*kwargs)
> 148 
> 149 def training\_step(self, batch, batch\_idx, optimizer\_idx=0):
> 
> ~/miniconda3/lib/python3.7/site-packages/torch/nn/modules/module.py in \_call\_impl(self, \*input, \*\*kwargs)
> 1100 if not (self.\_backward\_hooks or self.\_forward\_hooks or self.\_forward\_pre\_hooks or \_global\_backward\_hooks
> 1101 or \_global\_forward\_hooks or \_global\_forward\_pre\_hooks):
> \-\> 1102 return forward\_call(\*input, \*\*kwargs)
> 1103 # Do not call functions when jit is used
> 1104 full\_backward\_hooks, non\_full\_backward\_hooks = \[\], \[\]
> 
> ~/miniconda3/lib/python3.7/site-packages/scvi/module/base/\_decorators.py in auto\_transfer\_args(self, \*args, \*\*kwargs)
> 30 # decorator only necessary after training
> 31 if self.training:
> \---\> 32 return fn(self, \*args, \*\*kwargs)
> 33 
> 34 device = list(set(p.device for p in self.parameters()))
> 
> ~/miniconda3/lib/python3.7/site-packages/scvi/module/base/\_base\_module.py in forward(self, tensors, get\_inference\_input\_kwargs, get\_generative\_input\_kwargs, inference\_kwargs, generative\_kwargs, loss\_kwargs, compute\_loss)
> 143 tensors, \*\*get\_inference\_input\_kwargs
> 144 )
> \--\> 145 inference\_outputs = self.inference(\*\*inference\_inputs, \*\*inference\_kwargs)
> 146 generative\_inputs = self.\_get\_generative\_input(
> 147 tensors, inference\_outputs, \*\*get\_generative\_input\_kwargs
> 
> ~/miniconda3/lib/python3.7/site-packages/scvi/module/base/\_decorators.py in auto\_transfer\_args(self, \*args, \*\*kwargs)
> 30 # decorator only necessary after training
> 31 if self.training:
> \---\> 32 return fn(self, \*args, \*\*kwargs)
> 33 
> 34 device = list(set(p.device for p in self.parameters()))
> 
> ~/miniconda3/lib/python3.7/site-packages/scvi/module/\_multivae.py in inference(self, x, batch\_index, cont\_covs, cat\_covs, n\_samples)
> 293 # Z Encoders
> 294 qzm\_acc, qzv\_acc, z\_acc = self.z\_encoder\_accessibility(
> \--\> 295 encoder\_input\_accessibility, batch\_index, \*categorical\_input
> 296 )
> 297 qzm\_expr, qzv\_expr, z\_expr = self.z\_encoder\_expression(
> 
> ~/miniconda3/lib/python3.7/site-packages/torch/nn/modules/module.py in \_call\_impl(self, \*input, \*\*kwargs)
> 1100 if not (self.\_backward\_hooks or self.\_forward\_hooks or self.\_forward\_pre\_hooks or \_global\_backward\_hooks
> 1101 or \_global\_forward\_hooks or \_global\_forward\_pre\_hooks):
> \-\> 1102 return forward\_call(\*input, \*\*kwargs)
> 1103 # Do not call functions when jit is used
> 1104 full\_backward\_hooks, non\_full\_backward\_hooks = \[\], \[\]
> 
> ~/miniconda3/lib/python3.7/site-packages/scvi/nn/\_base\_components.py in forward(self, x, \*cat\_list)
> 292 q\_m = self.mean\_encoder(q)
> 293 q\_v = self.var\_activation(self.var\_encoder(q)) + self.var\_eps
> \--\> 294 latent = self.z\_transformation(reparameterize\_gaussian(q\_m, q\_v))
> 295 return q\_m, q\_v, latent
> 296 
> 
> ~/miniconda3/lib/python3.7/site-packages/scvi/nn/\_base\_components.py in reparameterize\_gaussian(mu, var)
> 11 
> 12 def reparameterize\_gaussian(mu, var):
> \---\> 13 return Normal(mu, var.sqrt()).rsample()
> 14 
> 15 
> 
> ~/miniconda3/lib/python3.7/site-packages/torch/distributions/normal.py in \_\_init\_\_(self, loc, scale, validate\_args)
> 48 else:
> 49 batch\_shape = self.loc.size()
> \---\> 50 super(Normal, self).\_\_init\_\_(batch\_shape, validate\_args=validate\_args)
> 51 
> 52 def expand(self, batch\_shape, \_instance=None):
> 
> ~/miniconda3/lib/python3.7/site-packages/torch/distributions/distribution.py in \_\_init\_\_(self, batch\_shape, event\_shape, validate\_args)
> 54 if not valid.all():
> 55 raise ValueError(
> \---\> 56 f"Expected parameter {param} "
> 57 f"({type(value).\_\_name\_\_} of shape {tuple(value.shape)}) "
> 58 f"of distribution {repr(self)} "
> 
> ValueError: Expected parameter loc (Tensor of shape (128, 19)) of distribution Normal(loc: torch.Size(\[128, 19\]), scale: torch.Size(\[128, 19\])) to satisfy the constraint Real(), but found invalid values:
> tensor(\[\[0.1109, 0.5803, 0.3902, ..., -0.4835, -0.8638, 0.0870\],
> \[0.7019, 0.4671, 0.4204, ..., 0.0102, -0.8212, 0.0126\],
> \[0.1285, 1.1512, -0.1905, ..., -0.4889, -0.0262, 0.0351\],
> ...,
> \[-0.0526, 0.3792, 0.6689, ..., 0.2085, 0.0496, 0.4914\],
> \[0.5148, 0.4604, 0.4606, ..., -0.0603, -0.3616, 0.4082\],
> \[0.5613, 0.6148, 0.5383, ..., -0.1725, -1.2356, -0.1374\]\],
> grad\_fn=\<AddmmBackward0\>)
> \`\`\`
> Can you help me with this?
> 
> Thanks in advance!
> Moritz
> 
> \#### Versions:
> 
> \> '0.14.5'
