Class: Gsplat::Ops::NormalizeDepth
- Inherits:
-
Autograd::Function
- Object
- Autograd::Function
- Gsplat::Ops::NormalizeDepth
- Defined in:
- lib/gsplat/ops/tensor_value_ops.rb
Overview
Divides one selected rendered depth channel by accumulated alpha.
Class Method Summary collapse
Methods inherited from Autograd::Function
Class Method Details
.backward(context, gradient) ⇒ Object
135 136 137 138 139 140 141 142 143 144 145 |
# File 'lib/gsplat/ops/tensor_value_ops.rb', line 135 def backward(context, gradient) alpha, depth, valid, depth_index = context.saved_values grad_rendered = gradient.dup grad_depth = gradient[*Array.new(gradient.ndim - 1, true).push(depth_index)] normalized_grad = gradient.class.zeros(*alpha.shape) normalized_grad[valid] = grad_depth[valid] / alpha[valid] if valid.any? grad_rendered[*Array.new(gradient.ndim - 1, true), depth_index] = normalized_grad grad_alpha = gradient.class.zeros(*alpha.shape) grad_alpha[valid] = -grad_depth[valid] * depth[valid] / (alpha[valid]**2) if valid.any? [grad_rendered, grad_alpha.reshape(*(alpha.shape + [1]))] end |
.forward(context, rendered, alphas, depth_index:) ⇒ Object
118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 |
# File 'lib/gsplat/ops/tensor_value_ops.rb', line 118 def forward(context, rendered, alphas, depth_index:) expected = rendered.shape[0...-1] + [1] unless alphas.shape == expected raise ShapeError, "expected alphas #{expected.inspect}, got #{alphas.shape.inspect}" end alpha = alphas[*Array.new(alphas.ndim - 1, true).push(0)] depth = rendered[*Array.new(rendered.ndim - 1, true).push(depth_index)] valid = alpha.gt(0) output = rendered.dup normalized = rendered.class.zeros(*alpha.shape) normalized[valid] = depth[valid] / alpha[valid] if valid.any? output[*Array.new(output.ndim - 1, true), depth_index] = normalized context.save(alpha, depth, valid, depth_index) output end |