Module: Gsplat::Training::Losses
- Defined in:
- lib/gsplat/training/losses.rb
Overview
Differentiable image losses and scalar image-quality metrics.
Defined Under Namespace
Classes: L1, Reconstruction, RegularizedReconstruction, StructuralSimilarity
Class Method Summary collapse
-
.l1(prediction, target) ⇒ Object
Returns mean absolute error, retaining an autograd graph when needed.
-
.psnr(prediction, target, max_value: 1.0) ⇒ Object
Returns peak signal-to-noise ratio in decibels.
-
.reconstruction(prediction, target, ssim_lambda: 0.2, layout: :auto) ⇒ Object
Returns
(1-lambda)*L1 + lambda*(1-SSIM). -
.regularized_reconstruction(prediction, target, opacities, scales, **options) ⇒ Object
Reconstruction objective plus optional MCMC opacity and scale regularizers.
-
.ssim(image_a, image_b, layout: :auto) ⇒ Object
Returns mean SSIM over batches, channels, and pixels.
Class Method Details
.l1(prediction, target) ⇒ Object
Returns mean absolute error, retaining an autograd graph when needed.
129 130 131 132 133 |
# File 'lib/gsplat/training/losses.rb', line 129 def l1(prediction, target) return L1.apply(prediction, target) if [prediction, target].any?(Autograd::Variable) L1.forward(Autograd::Context.new([false, false], [prediction, target]), prediction, target) end |
.psnr(prediction, target, max_value: 1.0) ⇒ Object
Returns peak signal-to-noise ratio in decibels.
176 177 178 179 180 181 182 183 184 185 186 187 188 |
# File 'lib/gsplat/training/losses.rb', line 176 def psnr(prediction, target, max_value: 1.0) prediction = Ops::TensorOps.data(prediction) target = Ops::TensorOps.data(target) unless prediction.shape == target.shape raise ShapeError, "PSNR shapes differ: #{prediction.shape.inspect} and #{target.shape.inspect}" end raise ArgumentError, "max_value must be positive" unless max_value.positive? mean_squared_error = ((prediction - target)**2).mean.to_f return Float::INFINITY if mean_squared_error.zero? 10 * ::Math.log10((max_value**2) / mean_squared_error) end |
.reconstruction(prediction, target, ssim_lambda: 0.2, layout: :auto) ⇒ Object
Returns (1-lambda)*L1 + lambda*(1-SSIM).
145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 |
# File 'lib/gsplat/training/losses.rb', line 145 def reconstruction(prediction, target, ssim_lambda: 0.2, layout: :auto) raise ArgumentError, "ssim_lambda must be between 0 and 1" unless ssim_lambda.between?(0.0, 1.0) if [prediction, target].any?(Autograd::Variable) return Reconstruction.apply( prediction, target, ssim_lambda: ssim_lambda, layout: layout ) end l1_value = (prediction - target).abs.mean.to_f ((1 - ssim_lambda) * l1_value) + (ssim_lambda * (1 - ssim(prediction, target, layout: layout))) end |
.regularized_reconstruction(prediction, target, opacities, scales, **options) ⇒ Object
Reconstruction objective plus optional MCMC opacity and scale regularizers.
162 163 164 165 166 167 168 169 170 171 172 173 |
# File 'lib/gsplat/training/losses.rb', line 162 def regularized_reconstruction(prediction, target, opacities, scales, **) RegularizedReconstruction.apply( prediction, target, opacities, scales, ssim_lambda: .fetch(:ssim_lambda, 0.2), opacity_reg: .fetch(:opacity_reg, 0.0), scale_reg: .fetch(:scale_reg, 0.0), layout: .fetch(:layout, :auto) ) end |
.ssim(image_a, image_b, layout: :auto) ⇒ Object
Returns mean SSIM over batches, channels, and pixels.
136 137 138 139 140 141 142 |
# File 'lib/gsplat/training/losses.rb', line 136 def ssim(image_a, image_b, layout: :auto) if [image_a, image_b].any?(Autograd::Variable) return StructuralSimilarity.apply(image_a, image_b, layout: layout) end Math::Ssim.forward(image_a, image_b, layout: layout).first end |