Class: Gsplat::Autograd::Function

Inherits:
Object
  • Object
show all
Defined in:
lib/gsplat/autograd/function.rb

Overview

Base class for differentiable coarse-grained operations.

Class Method Summary collapse

Class Method Details

.apply(*inputs) ⇒ Variable+

Executes forward and records a graph node when any input needs a gradient.

Parameters:

Returns:



40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
# File 'lib/gsplat/autograd/function.rb', line 40

def apply(*inputs, **)
  needs_input_grad = inputs.map { |input| input.is_a?(Variable) && input.requires_grad? }
  context = Context.new(needs_input_grad, inputs)
  raw_inputs = inputs.map { |input| input.is_a?(Variable) ? input.data : input }
  result = forward(context, *raw_inputs, **)
  multiple_outputs = result.is_a?(Array)
  output_data = multiple_outputs ? result : [result]
  requires_grad = Autograd.grad_enabled? && needs_input_grad.any?
  node = GraphNode.new(self, context, inputs) if requires_grad
  outputs = output_data.map do |data|
    validate_output!(data)
    Variable.new(data, requires_grad: requires_grad, creator: node)
  end
  node.outputs = outputs if node

  multiple_outputs ? outputs : outputs.first
end