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

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.

Raises:

  • (ArgumentError)


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).

Raises:

  • (ArgumentError)


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, **options)
  RegularizedReconstruction.apply(
    prediction,
    target,
    opacities,
    scales,
    ssim_lambda: options.fetch(:ssim_lambda, 0.2),
    opacity_reg: options.fetch(:opacity_reg, 0.0),
    scale_reg: options.fetch(:scale_reg, 0.0),
    layout: options.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