Class: GRX::NN::BatchNorm1d

Inherits:
Module
  • Object
show all
Defined in:
lib/grx/nn.rb

Overview

================================================================

BatchNorm1d — Normalizacion por batch

Instance Method Summary collapse

Methods inherited from Module

#call, #load_weights, #parameters, #save_weights, #zero_grad

Constructor Details

#initialize(num_features, epsilon: 1e-5, momentum: 0.1) ⇒ BatchNorm1d

Returns a new instance of BatchNorm1d.



321
322
323
324
325
326
327
328
329
330
331
332
# File 'lib/grx/nn.rb', line 321

def initialize(num_features, epsilon: 1e-5, momentum: 0.1)
  @num_features = num_features
  @epsilon      = epsilon
  @momentum     = momentum
  @training     = true

  @gamma = Tensor.ones([num_features],  requires_grad: true)
  @beta  = Tensor.zeros([num_features], requires_grad: true)

  @running_mean = Tensor.zeros([num_features])
  @running_var  = Tensor.ones([num_features])
end

Instance Method Details

#eval!Object



335
# File 'lib/grx/nn.rb', line 335

def eval!;  @training = false; self; end

#forward(x) ⇒ Object



337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
# File 'lib/grx/nn.rb', line 337

def forward(x)
  batch_size = x.shape[0]

  if @training
    batch_data = x.to_a
    means = Array.new(@num_features) do |j|
      batch_data.each_slice(@num_features).map { |row| row[j] }.sum / batch_size
    end
    vars = Array.new(@num_features) do |j|
      col = batch_data.each_slice(@num_features).map { |row| row[j] }
      col.sum { |v| (v - means[j]) ** 2 } / batch_size
    end

    means.each_with_index do |m, j|
      rm = @running_mean.to_a; rm[j] = (1 - @momentum) * rm[j] + @momentum * m
      @running_mean = Tensor.create(rm, [@num_features])
    end
    vars.each_with_index do |v, j|
      rv = @running_var.to_a; rv[j] = (1 - @momentum) * rv[j] + @momentum * v
      @running_var = Tensor.create(rv, [@num_features])
    end

    mean_t = Tensor.create(means, [@num_features])
    var_t  = Tensor.create(vars,  [@num_features])
  else
    mean_t = @running_mean
    var_t  = @running_var
  end

  norm_data = x.to_a.each_slice(@num_features).flat_map do |row|
    row.each_with_index.map do |v, j|
      x_hat = (v - mean_t.to_a[j]) / Math.sqrt(var_t.to_a[j] + @epsilon)
      @gamma.to_a[j] * x_hat + @beta.to_a[j]
    end
  end

  Tensor.create(norm_data, x.shape)
end

#to_sObject



376
# File 'lib/grx/nn.rb', line 376

def to_s = "BatchNorm1d(#{@num_features})"

#train!Object



334
# File 'lib/grx/nn.rb', line 334

def train!; @training = true;  self; end