Module: Gsplat::Math::Ssim

Defined in:
lib/gsplat/math/ssim.rb

Overview

SSIM forward and analytic VJP with a grouped Gaussian convolution.

Defined Under Namespace

Classes: Cache

Class Method Summary collapse

Class Method Details

.backward(cache, grad_output) ⇒ Object

rubocop:disable Metrics/AbcSize



44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
# File 'lib/gsplat/math/ssim.rb', line 44

def backward(cache, grad_output)
  factor = grad_output.to_f / cache.ssim_map.size
  denominator = cache.denominator_mean * cache.denominator_variance
  grad_num_mean = factor * cache.numerator_variance / denominator
  grad_num_variance = factor * cache.numerator_mean / denominator
  grad_den_mean = -factor * cache.ssim_map / cache.denominator_mean
  grad_den_variance = -factor * cache.ssim_map / cache.denominator_variance
  grad_mu_a = (2 * cache.mu_b * grad_num_mean) -
              (2 * cache.mu_b * grad_num_variance) +
              (2 * cache.mu_a * grad_den_mean) -
              (2 * cache.mu_a * grad_den_variance)
  grad_mu_b = (2 * cache.mu_a * grad_num_mean) -
              (2 * cache.mu_a * grad_num_variance) +
              (2 * cache.mu_b * grad_den_mean) -
              (2 * cache.mu_b * grad_den_variance)
  grad_cross = 2 * grad_num_variance
  filtered_variance = convolve(grad_den_variance, cache.kernel)
  filtered_cross = convolve(grad_cross, cache.kernel)
  grad_a = convolve(grad_mu_a, cache.kernel) +
           (2 * cache.image_a * filtered_variance) +
           (cache.image_b * filtered_cross)
  grad_b = convolve(grad_mu_b, cache.kernel) +
           (2 * cache.image_b * filtered_variance) +
           (cache.image_a * filtered_cross)
  [from_nchw(grad_a, cache.layout), from_nchw(grad_b, cache.layout)]
end

.forward(image_a, image_b, layout:) ⇒ Object

rubocop:disable Metrics/AbcSize



17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
# File 'lib/gsplat/math/ssim.rb', line 17

def forward(image_a, image_b, layout:)
  first, resolved_layout = to_nchw(image_a, layout)
  second, second_layout = to_nchw(image_b, layout)
  validate_inputs!(first, second, resolved_layout, second_layout)
  kernel = gaussian_kernel(first.class)
  mu_a = convolve(first, kernel)
  mu_b = convolve(second, kernel)
  variance_a = convolve(first**2, kernel) - (mu_a**2)
  variance_b = convolve(second**2, kernel) - (mu_b**2)
  covariance = convolve(first * second, kernel) - (mu_a * mu_b)
  numerator_mean = (2 * mu_a * mu_b) + (0.01**2)
  numerator_variance = (2 * covariance) + (0.03**2)
  denominator_mean = (mu_a**2) + (mu_b**2) + (0.01**2)
  denominator_variance = variance_a + variance_b + (0.03**2)
  ssim_map = (numerator_mean * numerator_variance) /
             (denominator_mean * denominator_variance)
  cache = Cache.new(
    image_a: first, image_b: second, mu_a: mu_a, mu_b: mu_b, ssim_map: ssim_map,
    numerator_mean: numerator_mean, numerator_variance: numerator_variance,
    denominator_mean: denominator_mean, denominator_variance: denominator_variance,
    kernel: kernel, layout: resolved_layout
  )
  [ssim_map.mean.to_f, cache]
end