Class: Gsplat::Autograd::Function
- Inherits:
-
Object
- Object
- Gsplat::Autograd::Function
- Defined in:
- lib/gsplat/autograd/function.rb
Overview
Base class for differentiable coarse-grained operations.
Direct Known Subclasses
Ops::Accumulate, Ops::AddClampMin, Ops::CameraBroadcast, Ops::CameraDirections, Ops::ConcatCoefficients, Ops::ConcatDepth, Ops::ConcatFeatures, Ops::DepthFeatures, Ops::Eval3dRasterize, Ops::Exp, Ops::FeatureSlice, Ops::FullyFusedProjection, Ops::Multiply, Ops::NormalizeDepth, Ops::NormalizeQuaternion, Ops::QuatScaleToCovarPreci, Ops::RasterizeToPixels, Ops::Sigmoid, Ops::SphericalHarmonics, Training::ImageFitter::MeanSquaredError, Training::Losses::L1, Training::Losses::Reconstruction, Training::Losses::RegularizedReconstruction, Training::Losses::StructuralSimilarity
Class Method Summary collapse
-
.apply(*inputs) ⇒ Variable+
Executes forward and records a graph node when any input needs a gradient.
Class Method Details
.apply(*inputs) ⇒ Variable+
Executes forward and records a graph node when any input needs a gradient.
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 |