Class: Gsplat::Ops::NormalizeDepth

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

apply

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