Class: Daimond::NN::MaxPool2dRust
- Defined in:
- lib/daimond/nn/max_pool2d_rust.rb
Instance Method Summary collapse
- #forward(input) ⇒ Object
-
#initialize(kernel_size) ⇒ MaxPool2dRust
constructor
A new instance of MaxPool2dRust.
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 |