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
-
.backward(cache, grad_output) ⇒ Object
rubocop:disable Metrics/AbcSize.
-
.forward(image_a, image_b, layout:) ⇒ Object
rubocop:disable Metrics/AbcSize.
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 |