Class: Daimond::Loss::CrossEntropyLoss

Inherits:
Object
  • Object
show all
Defined in:
lib/daimond/loss/cross_entropy.rb

Instance Method Summary collapse

Constructor Details

#initialize ⇒ CrossEntropyLoss

Returns a new instance of CrossEntropyLoss.



7
8
# File 'lib/daimond/loss/cross_entropy.rb', line 7

def initialize
end

Instance Method Details

#call(pred, target) ⇒ Object



40
41
42
# File 'lib/daimond/loss/cross_entropy.rb', line 40

def call(pred, target)
  forward(pred, target)
end

#forward(pred, target) ⇒ Object



10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
# File 'lib/daimond/loss/cross_entropy.rb', line 10

def forward(pred, target)
  # pred: [batch_size, 10] после softmax
  # target: [batch_size] метки классов (0-9)
  batch_size = pred.shape[0]

  # Вычисляем loss для мониторинга (не используется в backward)
  log_probs = Numo::NMath.log(pred.data)
  correct_log_probs = Numo::DFloat.zeros(batch_size)

  batch_size.times do |i|
    correct_log_probs[i] = log_probs[i, target.data[i]]
  end

  loss_value = -correct_log_probs.mean

  out = Tensor.new(Numo::DFloat[loss_value], prev: [pred], op: 'cross_entropy')

  # Backward: gradient of cross_entropy + softmax = pred - one_hot(target)
  out._backward = lambda do
    grad_input = pred.data.dup  # softmax output
    batch_size.times do |i|
      grad_input[i, target.data[i]] -= 1.0
    end
    grad_input /= batch_size
    pred.grad += grad_input
  end

  out
end