Class: Gsplat::Ops::QuatScaleToCovarPreci
- Inherits:
-
Autograd::Function
- Object
- Autograd::Function
- Gsplat::Ops::QuatScaleToCovarPreci
- Defined in:
- lib/gsplat/ops/quat_scale_to_covar_preci.rb
Overview
Differentiable conversion from wxyz quaternions/scales to covariance and precision.
Class Method Summary collapse
-
.apply(quaternions, scales, compute_covar: true, compute_preci: true, triu: false) ⇒ Array<(Autograd::Variable, nil)>
Applies the operation while retaining Python gsplat's two tuple positions.
-
.backward(context, *grad_outputs) ⇒ Object
private
Propagates matrix gradients to quaternions and scales.
-
.forward(context, quaternions, scales, **options) ⇒ Object
private
Evaluates covariance/precision tensors and records inputs for VJP.
Class Method Details
.apply(quaternions, scales, compute_covar: true, compute_preci: true, triu: false) ⇒ Array<(Autograd::Variable, nil)>
Applies the operation while retaining Python gsplat's two tuple positions.
19 20 21 22 23 24 25 26 |
# File 'lib/gsplat/ops/quat_scale_to_covar_preci.rb', line 19 def apply(quaternions, scales, compute_covar: true, compute_preci: true, triu: false) validate_selection!(compute_covar, compute_preci) outputs = super return outputs if compute_covar && compute_preci return [outputs, nil] if compute_covar [nil, outputs] end |
.backward(context, *grad_outputs) ⇒ Object
This method is part of a private API. You should avoid using this method if possible, as it may be removed or be changed in the future.
Propagates matrix gradients to quaternions and scales.
50 51 52 53 54 55 56 57 58 59 60 61 62 |
# File 'lib/gsplat/ops/quat_scale_to_covar_preci.rb', line 50 def backward(context, *grad_outputs) quaternions, scales, compute_covar, compute_preci, triu = context.saved_values grad_covar = compute_covar ? grad_outputs.shift : nil grad_preci = compute_preci ? grad_outputs.shift : nil Backend.dispatch( :quat_scale_to_covar_preci_backward, quaternions, scales, grad_covar, grad_preci, triu: triu ) end |
.forward(context, quaternions, scales, **options) ⇒ Object
This method is part of a private API. You should avoid using this method if possible, as it may be removed or be changed in the future.
Evaluates covariance/precision tensors and records inputs for VJP.
30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 |
# File 'lib/gsplat/ops/quat_scale_to_covar_preci.rb', line 30 def forward(context, quaternions, scales, **) compute_covar = .fetch(:compute_covar) compute_preci = .fetch(:compute_preci) triu = .fetch(:triu) context.save(quaternions, scales, compute_covar, compute_preci, triu) covariance, precision = Backend.dispatch( :quat_scale_to_covar_preci_forward, quaternions, scales, compute_covar: compute_covar, compute_preci: compute_preci, triu: triu ) return [covariance, precision] if compute_covar && compute_preci compute_covar ? covariance : precision end |