Class: Daimond::NN::MaxPool2dRust

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

Instance Method Summary collapse

Methods inherited from Module

#call, #load, #parameters, #save, #zero_grad

Constructor Details

#initialize(kernel_size) ⇒ MaxPool2dRust

Returns a new instance of MaxPool2dRust.



6
7
8
9
# File 'lib/daimond/nn/max_pool2d_rust.rb', line 6

def initialize(kernel_size)
  super()
  @kernel_size = kernel_size
end

Instance Method Details

#forward(input) ⇒ Object



11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
# File 'lib/daimond/nn/max_pool2d_rust.rb', line 11

def forward(input)

  batch = input.shape[0]
  channels = input.shape[1]
  h = input.shape[2]
  w = input.shape[3]
  k = @kernel_size

  if Daimond::RustBackend.available?
    output_data = Daimond::RustBackend.maxpool2d(
      input.data, batch, channels, h, w, k
    )

    out = Tensor.new(output_data, prev: [input], op: 'maxpool2d_rust')
    out._backward = lambda {}
    return out
  else
    raise "Rust backend required for MaxPool2dRust"
  end
end