Module: Daimond::RustBackend

Defined in:
lib/daimond/rust_bridge.rb,
lib/daimond/rust_backend.rb

Overview

Модуль-обертка для вызовов

Class Method Summary collapse

Class Method Details

.available?Boolean

Проверка доступности

Returns:

  • (Boolean)


10
11
12
# File 'lib/daimond/rust_backend.rb', line 10

def available?
  Daimond.rust_available?
end

.conv2d(input_data, weight_data, bias_data, batch, in_c, out_c, h, w, k) ⇒ Object



16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
# File 'lib/daimond/rust_bridge.rb', line 16

def conv2d(input_data, weight_data, bias_data, batch, in_c, out_c, h, w, k)
  return nil unless available?

  flat_input = input_data.flatten.to_a
  flat_weight = weight_data.flatten.to_a
  flat_bias = bias_data.to_a

  result_flat = Daimond::Rust.conv2d_native(
    flat_input, flat_weight, flat_bias,
    batch, in_c, out_c, h, w, k
  )

  h_out = h - k + 1
  w_out = w - k + 1
  Numo::DFloat[*result_flat].reshape(batch, out_c, h_out, w_out)
end

.matmul(a, b) ⇒ Object

Обертка для матричного умножения



17
18
19
20
21
# File 'lib/daimond/rust_backend.rb', line 17

def self.matmul(a, b)
  # Здесь будет код конвертации Ruby -> Rust -> Ruby
  # Пока просто возвращаем Rust тензор
  Rust::Tensor.zeros(a.shape[0], b.shape[1])
end

.matmul_data(narray_a, narray_b) ⇒ Object



46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
# File 'lib/daimond/rust_bridge.rb', line 46

def matmul_data(narray_a, narray_b)
  return nil unless available?

  shape_a = narray_a.shape
  shape_b = narray_b.shape

  flat_a = narray_a.flatten.to_a
  flat_b = narray_b.flatten.to_a

  result_flat = Daimond::Rust.fast_matmul_flat(
    flat_a, flat_b, shape_a[0], shape_a[1], shape_b[1]
  )

  Numo::DFloat[*result_flat].reshape(shape_a[0], shape_b[1])
end

.maxpool2d(input_data, batch, channels, h, w, k) ⇒ Object



33
34
35
36
37
38
39
40
41
42
43
44
# File 'lib/daimond/rust_bridge.rb', line 33

def maxpool2d(input_data, batch, channels, h, w, k)
  return nil unless available?

  flat_input = input_data.flatten.to_a
  result_flat = Daimond::Rust.maxpool2d_native(
    flat_input, batch, channels, h, w, k
  )

  h_out = h / k
  w_out = w / k
  Numo::DFloat[*result_flat].reshape(batch, channels, h_out, w_out)
end