Class: Gsplat::Ops::QuatScaleToCovarPreci

Inherits:
Autograd::Function show all
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

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.

Parameters:

  • quaternions (Autograd::Variable, Numo::NArray)

    [...,4]

  • scales (Autograd::Variable, Numo::NArray)

    [...,3]

  • compute_covar (Boolean) (defaults to: true)
  • compute_preci (Boolean) (defaults to: true)
  • triu (Boolean) (defaults to: false)

    return [...,6] upper triangles when true

Returns:



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, **options)
  compute_covar = options.fetch(:compute_covar)
  compute_preci = options.fetch(:compute_preci)
  triu = options.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