grx-tensor 0.2.0 → 0.2.1

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
data/ext/grx/grx_core.c CHANGED
@@ -1,18 +1,20 @@
1
1
  /*
2
- * grx_core.c — Núcleo C de GRX
3
- * =============================================================
4
- * Optimizaciones activas:
5
- * - AVX2 + FMA: 4 doubles/ciclo con multiply-add fusionado
6
- * - Loop unrolling x2: mayor ILP (Instruction Level Parallelism)
7
- * - restrict: elimina alias analysis, permite más vectorización auto
8
- * - Memoria alineada 32 bytes: habilita _mm256_load_pd (más rápido que loadu)
9
- * - matmul con tiling: respeta líneas de caché L1 (64 bytes = 8 doubles)
10
- * - Adam con FMA: beta*m + (1-beta)*grad en una pasada
11
- * =============================================================
2
+ * grx_core.c — Nucleo C de GRX con Despacho Dinamico Multi-Target SIMD
3
+ * ====================================================================
4
+ * Caracteristicas de compatibilidad y rendimiento:
5
+ * - Compilacion base universal (sin banderas forzadas globales -mavx2)
6
+ * - Despacho en tiempo de ejecucion segun capacidades reales de la CPU:
7
+ * * AVX2 + FMA: 4 doubles/ciclo con multiply-add fusionado
8
+ * * Fallback escalar seguro: compatible con 100% de CPUs (x86_64, ARM, VMs)
9
+ * - Memoria alineada a 32 bytes: posix_memalign en Unix, _aligned_malloc en Windows
10
+ * - Matmul con cache tiling (L1 64 bytes)
11
+ * - Optimizadores in-place (SGD, Adam con correccion de sesgo)
12
+ * - Generadores de pesos xorshift64 y Box-Muller universales
13
+ * ====================================================================
12
14
  */
13
15
 
14
- #define _USE_MATH_DEFINES /* M_PI en Windows/MSVC */
15
- #define _POSIX_C_SOURCE 200809L /* posix_memalign, M_PI en glibc */
16
+ #define _USE_MATH_DEFINES
17
+ #define _POSIX_C_SOURCE 200809L
16
18
  #include "grx_core.h"
17
19
  #include <stdlib.h>
18
20
  #include <stdint.h>
@@ -25,36 +27,65 @@
25
27
  #define M_PI 3.14159265358979323846
26
28
  #endif
27
29
 
28
- #if defined(__AVX2__) && defined(__FMA__)
30
+ #if defined(__GNUC__) || defined(__clang__)
31
+ #define GRX_TARGET_AVX2 __attribute__((target("avx2,fma")))
29
32
  #include <immintrin.h>
30
- #define GRX_AVX2_FMA 1
31
- #elif defined(__AVX2__)
32
- #include <immintrin.h>
33
- #define GRX_AVX2 1
34
- #elif defined(__SSE2__)
35
- #include <emmintrin.h>
36
- #define GRX_SSE2 1
33
+ #else
34
+ #define GRX_TARGET_AVX2
37
35
  #endif
38
36
 
39
37
  #define TILE 8
40
38
 
41
39
  /* ================================================================
42
- * MEMORIA
40
+ * DETECCION DE CPU Y NIVEL SIMD
41
+ * ================================================================ */
42
+
43
+ static int g_simd_level = -1;
44
+
45
+ static int grx_detect_simd(void) {
46
+ if (__builtin_expect(g_simd_level != -1, 1)) {
47
+ return g_simd_level;
48
+ }
49
+
50
+ #if (defined(__x86_64__) || defined(_M_X64) || defined(__i386__) || defined(_M_IX86)) && (defined(__GNUC__) || defined(__clang__))
51
+ __builtin_cpu_init();
52
+ if (__builtin_cpu_supports("avx2") && __builtin_cpu_supports("fma")) {
53
+ g_simd_level = 2; /* AVX2 + FMA */
54
+ return 2;
55
+ } else if (__builtin_cpu_supports("sse2")) {
56
+ g_simd_level = 1; /* SSE2 */
57
+ return 1;
58
+ }
59
+ #endif
60
+
61
+ g_simd_level = 0; /* Escalar universal */
62
+ return 0;
63
+ }
64
+
65
+ GRX_API int grx_simd_level(void) {
66
+ return grx_detect_simd();
67
+ }
68
+
69
+ /* ================================================================
70
+ * GESTION DE MEMORIA
43
71
  * ================================================================ */
44
72
 
45
73
  GRX_API double* grx_alloc(size_t n) {
46
74
  if (__builtin_expect(n == 0, 0)) return NULL;
47
75
  void *ptr = NULL;
48
- #if defined(_WIN32)
76
+ #if defined(_WIN32) || defined(_WIN64)
49
77
  ptr = _aligned_malloc(n * sizeof(double), 32);
50
78
  #else
51
- if (posix_memalign(&ptr, 32, n * sizeof(double)) != 0) return NULL;
79
+ if (posix_memalign(&ptr, 32, n * sizeof(double)) != 0) {
80
+ ptr = malloc(n * sizeof(double));
81
+ }
52
82
  #endif
53
83
  return (double*)ptr;
54
84
  }
55
85
 
56
86
  GRX_API void grx_free(double *ptr) {
57
- #if defined(_WIN32)
87
+ if (!ptr) return;
88
+ #if defined(_WIN32) || defined(_WIN64)
58
89
  _aligned_free(ptr);
59
90
  #else
60
91
  free(ptr);
@@ -62,268 +93,370 @@ GRX_API void grx_free(double *ptr) {
62
93
  }
63
94
 
64
95
  /* ================================================================
65
- * MACROS SIMD INTERNOS
96
+ * ARITMETICA ELEMENT-WISE: KERNELS AVX2 Y ESCALARES
66
97
  * ================================================================ */
67
98
 
68
- /* Carga/store: usa aligned si tenemos AVX2+FMA (memoria siempre alineada a 32b) */
69
- #ifdef GRX_AVX2_FMA
70
- #define VLD(p) _mm256_load_pd(p)
71
- #define VST(p, v) _mm256_store_pd(p, v)
72
- #elif defined(GRX_AVX2)
73
- #define VLD(p) _mm256_loadu_pd(p)
74
- #define VST(p, v) _mm256_storeu_pd(p, v)
75
- #endif
76
-
77
- /* ================================================================
78
- * ELEMENT-WISE ARITMÉTICA
79
- * ================================================================ */
99
+ #if defined(__GNUC__) || defined(__clang__)
80
100
 
81
- #define BINOP_BODY(op_avx, op_scalar) \
82
- size_t i = 0; \
83
- for (; i + 8 <= n; i += 8) { \
84
- VST(out+i, op_avx(VLD(a+i), VLD(b+i))); \
85
- VST(out+i+4, op_avx(VLD(a+i+4), VLD(b+i+4))); \
86
- } \
87
- for (; i + 4 <= n; i += 4) VST(out+i, op_avx(VLD(a+i), VLD(b+i)));\
88
- for (; i < n; i++) out[i] = op_scalar(a[i], b[i]);
89
-
90
- #define SCALAR_ADD(x,y) ((x)+(y))
91
- #define SCALAR_SUB(x,y) ((x)-(y))
92
- #define SCALAR_MUL(x,y) ((x)*(y))
93
- #define SCALAR_DIV(x,y) ((x)/(y))
94
-
95
- GRX_API void grx_add(const double * restrict a, const double * restrict b,
96
- double * restrict out, size_t n) {
97
- #if defined(GRX_AVX2_FMA) || defined(GRX_AVX2)
98
- BINOP_BODY(_mm256_add_pd, SCALAR_ADD)
99
- #else
100
- for (size_t i = 0; i < n; i++) out[i] = a[i] + b[i];
101
- #endif
101
+ GRX_TARGET_AVX2
102
+ static void grx_add_avx2(const double *a, const double *b, double *out, size_t n) {
103
+ size_t i = 0;
104
+ for (; i + 8 <= n; i += 8) {
105
+ _mm256_storeu_pd(out + i, _mm256_add_pd(_mm256_loadu_pd(a + i), _mm256_loadu_pd(b + i)));
106
+ _mm256_storeu_pd(out + i + 4, _mm256_add_pd(_mm256_loadu_pd(a + i + 4), _mm256_loadu_pd(b + i + 4)));
107
+ }
108
+ for (; i + 4 <= n; i += 4) {
109
+ _mm256_storeu_pd(out + i, _mm256_add_pd(_mm256_loadu_pd(a + i), _mm256_loadu_pd(b + i)));
110
+ }
111
+ for (; i < n; i++) out[i] = a[i] + b[i];
102
112
  }
103
113
 
104
- GRX_API void grx_sub(const double * restrict a, const double * restrict b,
105
- double * restrict out, size_t n) {
106
- #if defined(GRX_AVX2_FMA) || defined(GRX_AVX2)
107
- BINOP_BODY(_mm256_sub_pd, SCALAR_SUB)
108
- #else
109
- for (size_t i = 0; i < n; i++) out[i] = a[i] - b[i];
110
- #endif
114
+ GRX_TARGET_AVX2
115
+ static void grx_sub_avx2(const double *a, const double *b, double *out, size_t n) {
116
+ size_t i = 0;
117
+ for (; i + 8 <= n; i += 8) {
118
+ _mm256_storeu_pd(out + i, _mm256_sub_pd(_mm256_loadu_pd(a + i), _mm256_loadu_pd(b + i)));
119
+ _mm256_storeu_pd(out + i + 4, _mm256_sub_pd(_mm256_loadu_pd(a + i + 4), _mm256_loadu_pd(b + i + 4)));
120
+ }
121
+ for (; i + 4 <= n; i += 4) {
122
+ _mm256_storeu_pd(out + i, _mm256_sub_pd(_mm256_loadu_pd(a + i), _mm256_loadu_pd(b + i)));
123
+ }
124
+ for (; i < n; i++) out[i] = a[i] - b[i];
111
125
  }
112
126
 
113
- GRX_API void grx_mul(const double * restrict a, const double * restrict b,
114
- double * restrict out, size_t n) {
115
- #if defined(GRX_AVX2_FMA) || defined(GRX_AVX2)
116
- BINOP_BODY(_mm256_mul_pd, SCALAR_MUL)
117
- #else
118
- for (size_t i = 0; i < n; i++) out[i] = a[i] * b[i];
119
- #endif
127
+ GRX_TARGET_AVX2
128
+ static void grx_mul_avx2(const double *a, const double *b, double *out, size_t n) {
129
+ size_t i = 0;
130
+ for (; i + 8 <= n; i += 8) {
131
+ _mm256_storeu_pd(out + i, _mm256_mul_pd(_mm256_loadu_pd(a + i), _mm256_loadu_pd(b + i)));
132
+ _mm256_storeu_pd(out + i + 4, _mm256_mul_pd(_mm256_loadu_pd(a + i + 4), _mm256_loadu_pd(b + i + 4)));
133
+ }
134
+ for (; i + 4 <= n; i += 4) {
135
+ _mm256_storeu_pd(out + i, _mm256_mul_pd(_mm256_loadu_pd(a + i), _mm256_loadu_pd(b + i)));
136
+ }
137
+ for (; i < n; i++) out[i] = a[i] * b[i];
120
138
  }
121
139
 
122
- GRX_API void grx_div(const double * restrict a, const double * restrict b,
123
- double * restrict out, size_t n) {
124
- #if defined(GRX_AVX2_FMA) || defined(GRX_AVX2)
125
- BINOP_BODY(_mm256_div_pd, SCALAR_DIV)
126
- #else
127
- for (size_t i = 0; i < n; i++) out[i] = a[i] / b[i];
128
- #endif
140
+ GRX_TARGET_AVX2
141
+ static void grx_div_avx2(const double *a, const double *b, double *out, size_t n) {
142
+ size_t i = 0;
143
+ for (; i + 8 <= n; i += 8) {
144
+ _mm256_storeu_pd(out + i, _mm256_div_pd(_mm256_loadu_pd(a + i), _mm256_loadu_pd(b + i)));
145
+ _mm256_storeu_pd(out + i + 4, _mm256_div_pd(_mm256_loadu_pd(a + i + 4), _mm256_loadu_pd(b + i + 4)));
146
+ }
147
+ for (; i + 4 <= n; i += 4) {
148
+ _mm256_storeu_pd(out + i, _mm256_div_pd(_mm256_loadu_pd(a + i), _mm256_loadu_pd(b + i)));
149
+ }
150
+ for (; i < n; i++) out[i] = a[i] / b[i];
129
151
  }
130
152
 
131
- GRX_API void grx_scale(const double * restrict a, double s,
132
- double * restrict out, size_t n) {
133
- #if defined(GRX_AVX2_FMA) || defined(GRX_AVX2)
153
+ GRX_TARGET_AVX2
154
+ static void grx_scale_avx2(const double *a, double s, double *out, size_t n) {
134
155
  __m256d vs = _mm256_set1_pd(s);
135
156
  size_t i = 0;
136
157
  for (; i + 8 <= n; i += 8) {
137
- VST(out+i, _mm256_mul_pd(VLD(a+i), vs));
138
- VST(out+i+4, _mm256_mul_pd(VLD(a+i+4), vs));
158
+ _mm256_storeu_pd(out + i, _mm256_mul_pd(_mm256_loadu_pd(a + i), vs));
159
+ _mm256_storeu_pd(out + i + 4, _mm256_mul_pd(_mm256_loadu_pd(a + i + 4), vs));
160
+ }
161
+ for (; i + 4 <= n; i += 4) {
162
+ _mm256_storeu_pd(out + i, _mm256_mul_pd(_mm256_loadu_pd(a + i), vs));
139
163
  }
140
- for (; i + 4 <= n; i += 4) VST(out+i, _mm256_mul_pd(VLD(a+i), vs));
141
164
  for (; i < n; i++) out[i] = a[i] * s;
142
- #else
143
- for (size_t i = 0; i < n; i++) out[i] = a[i] * s;
144
- #endif
145
165
  }
146
166
 
147
- GRX_API void grx_add_scalar(const double * restrict a, double s,
148
- double * restrict out, size_t n) {
149
- #if defined(GRX_AVX2_FMA) || defined(GRX_AVX2)
167
+ GRX_TARGET_AVX2
168
+ static void grx_add_scalar_avx2(const double *a, double s, double *out, size_t n) {
150
169
  __m256d vs = _mm256_set1_pd(s);
151
170
  size_t i = 0;
152
171
  for (; i + 8 <= n; i += 8) {
153
- VST(out+i, _mm256_add_pd(VLD(a+i), vs));
154
- VST(out+i+4, _mm256_add_pd(VLD(a+i+4), vs));
172
+ _mm256_storeu_pd(out + i, _mm256_add_pd(_mm256_loadu_pd(a + i), vs));
173
+ _mm256_storeu_pd(out + i + 4, _mm256_add_pd(_mm256_loadu_pd(a + i + 4), vs));
174
+ }
175
+ for (; i + 4 <= n; i += 4) {
176
+ _mm256_storeu_pd(out + i, _mm256_add_pd(_mm256_loadu_pd(a + i), vs));
155
177
  }
156
- for (; i + 4 <= n; i += 4) VST(out+i, _mm256_add_pd(VLD(a+i), vs));
157
178
  for (; i < n; i++) out[i] = a[i] + s;
158
- #else
159
- for (size_t i = 0; i < n; i++) out[i] = a[i] + s;
179
+ }
180
+
181
+ GRX_TARGET_AVX2
182
+ static double grx_sum_avx2(const double *a, size_t n) {
183
+ __m256d v0 = _mm256_setzero_pd(), v1 = _mm256_setzero_pd();
184
+ size_t i = 0;
185
+ for (; i + 8 <= n; i += 8) {
186
+ v0 = _mm256_add_pd(v0, _mm256_loadu_pd(a + i));
187
+ v1 = _mm256_add_pd(v1, _mm256_loadu_pd(a + i + 4));
188
+ }
189
+ v0 = _mm256_add_pd(v0, v1);
190
+ for (; i + 4 <= n; i += 4) {
191
+ v0 = _mm256_add_pd(v0, _mm256_loadu_pd(a + i));
192
+ }
193
+ double tmp[4];
194
+ _mm256_storeu_pd(tmp, v0);
195
+ double acc = tmp[0] + tmp[1] + tmp[2] + tmp[3];
196
+ for (; i < n; i++) acc += a[i];
197
+ return acc;
198
+ }
199
+
200
+ GRX_TARGET_AVX2
201
+ static double grx_dot_avx2(const double *a, const double *b, size_t n) {
202
+ __m256d acc0 = _mm256_setzero_pd(), acc1 = _mm256_setzero_pd();
203
+ size_t i = 0;
204
+ for (; i + 8 <= n; i += 8) {
205
+ acc0 = _mm256_fmadd_pd(_mm256_loadu_pd(a + i), _mm256_loadu_pd(b + i), acc0);
206
+ acc1 = _mm256_fmadd_pd(_mm256_loadu_pd(a + i + 4), _mm256_loadu_pd(b + i + 4), acc1);
207
+ }
208
+ acc0 = _mm256_add_pd(acc0, acc1);
209
+ for (; i + 4 <= n; i += 4) {
210
+ acc0 = _mm256_fmadd_pd(_mm256_loadu_pd(a + i), _mm256_loadu_pd(b + i), acc0);
211
+ }
212
+ double tmp[4];
213
+ _mm256_storeu_pd(tmp, acc0);
214
+ double sum = tmp[0] + tmp[1] + tmp[2] + tmp[3];
215
+ for (; i < n; i++) sum += a[i] * b[i];
216
+ return sum;
217
+ }
218
+
219
+ GRX_TARGET_AVX2
220
+ static void grx_relu_avx2(const double *a, double *out, size_t n) {
221
+ __m256d zero = _mm256_setzero_pd();
222
+ size_t i = 0;
223
+ for (; i + 8 <= n; i += 8) {
224
+ _mm256_storeu_pd(out + i, _mm256_max_pd(_mm256_loadu_pd(a + i), zero));
225
+ _mm256_storeu_pd(out + i + 4, _mm256_max_pd(_mm256_loadu_pd(a + i + 4), zero));
226
+ }
227
+ for (; i + 4 <= n; i += 4) {
228
+ _mm256_storeu_pd(out + i, _mm256_max_pd(_mm256_loadu_pd(a + i), zero));
229
+ }
230
+ for (; i < n; i++) out[i] = a[i] > 0.0 ? a[i] : 0.0;
231
+ }
232
+
233
+ GRX_TARGET_AVX2
234
+ static void grx_adam_step_avx2(double *param, double *m, double *v,
235
+ const double *grad, double lr,
236
+ double beta1, double beta2, double epsilon,
237
+ double beta1t, double beta2t, size_t n) {
238
+ __m256d vb1 = _mm256_set1_pd(beta1);
239
+ __m256d vb2 = _mm256_set1_pd(beta2);
240
+ __m256d vom1 = _mm256_set1_pd(1.0 - beta1);
241
+ __m256d vom2 = _mm256_set1_pd(1.0 - beta2);
242
+ __m256d vcb1 = _mm256_set1_pd(1.0 / (1.0 - beta1t));
243
+ __m256d vcb2 = _mm256_set1_pd(1.0 / (1.0 - beta2t));
244
+ __m256d vlr = _mm256_set1_pd(lr);
245
+ __m256d veps = _mm256_set1_pd(epsilon);
246
+
247
+ size_t i = 0;
248
+ for (; i + 4 <= n; i += 4) {
249
+ __m256d g = _mm256_loadu_pd(grad + i);
250
+ __m256d mi = _mm256_loadu_pd(m + i);
251
+ __m256d vi = _mm256_loadu_pd(v + i);
252
+ __m256d p = _mm256_loadu_pd(param + i);
253
+
254
+ mi = _mm256_fmadd_pd(vb1, mi, _mm256_mul_pd(vom1, g));
255
+ vi = _mm256_fmadd_pd(vb2, vi, _mm256_mul_pd(vom2, _mm256_mul_pd(g, g)));
256
+
257
+ _mm256_storeu_pd(m + i, mi);
258
+ _mm256_storeu_pd(v + i, vi);
259
+
260
+ __m256d m_hat = _mm256_mul_pd(mi, vcb1);
261
+ __m256d v_hat = _mm256_mul_pd(vi, vcb2);
262
+ __m256d denom = _mm256_add_pd(_mm256_sqrt_pd(v_hat), veps);
263
+ __m256d step = _mm256_div_pd(_mm256_mul_pd(vlr, m_hat), denom);
264
+
265
+ _mm256_storeu_pd(param + i, _mm256_sub_pd(p, step));
266
+ }
267
+
268
+ double cb1 = 1.0 / (1.0 - beta1t);
269
+ double cb2 = 1.0 / (1.0 - beta2t);
270
+ for (; i < n; i++) {
271
+ double g = grad[i];
272
+ m[i] = beta1 * m[i] + (1.0 - beta1) * g;
273
+ v[i] = beta2 * v[i] + (1.0 - beta2) * g * g;
274
+ double m_hat = m[i] * cb1;
275
+ double v_hat = v[i] * cb2;
276
+ param[i] -= lr * m_hat / (sqrt(v_hat) + epsilon);
277
+ }
278
+ }
279
+
280
+ #endif /* GCC / Clang */
281
+
282
+ /* ================================================================
283
+ * FUNCIONES PUBLICAS CON DESPACHO DINAMICO
284
+ * ================================================================ */
285
+
286
+ GRX_API void grx_add(const double *a, const double *b, double *out, size_t n) {
287
+ #if defined(__GNUC__) || defined(__clang__)
288
+ if (grx_detect_simd() >= 2) {
289
+ grx_add_avx2(a, b, out, n);
290
+ return;
291
+ }
292
+ #endif
293
+ for (size_t i = 0; i < n; i++) out[i] = a[i] + b[i];
294
+ }
295
+
296
+ GRX_API void grx_sub(const double *a, const double *b, double *out, size_t n) {
297
+ #if defined(__GNUC__) || defined(__clang__)
298
+ if (grx_detect_simd() >= 2) {
299
+ grx_sub_avx2(a, b, out, n);
300
+ return;
301
+ }
302
+ #endif
303
+ for (size_t i = 0; i < n; i++) out[i] = a[i] - b[i];
304
+ }
305
+
306
+ GRX_API void grx_mul(const double *a, const double *b, double *out, size_t n) {
307
+ #if defined(__GNUC__) || defined(__clang__)
308
+ if (grx_detect_simd() >= 2) {
309
+ grx_mul_avx2(a, b, out, n);
310
+ return;
311
+ }
312
+ #endif
313
+ for (size_t i = 0; i < n; i++) out[i] = a[i] * b[i];
314
+ }
315
+
316
+ GRX_API void grx_div(const double *a, const double *b, double *out, size_t n) {
317
+ #if defined(__GNUC__) || defined(__clang__)
318
+ if (grx_detect_simd() >= 2) {
319
+ grx_div_avx2(a, b, out, n);
320
+ return;
321
+ }
322
+ #endif
323
+ for (size_t i = 0; i < n; i++) out[i] = a[i] / b[i];
324
+ }
325
+
326
+ GRX_API void grx_scale(const double *a, double s, double *out, size_t n) {
327
+ #if defined(__GNUC__) || defined(__clang__)
328
+ if (grx_detect_simd() >= 2) {
329
+ grx_scale_avx2(a, s, out, n);
330
+ return;
331
+ }
332
+ #endif
333
+ for (size_t i = 0; i < n; i++) out[i] = a[i] * s;
334
+ }
335
+
336
+ GRX_API void grx_add_scalar(const double *a, double s, double *out, size_t n) {
337
+ #if defined(__GNUC__) || defined(__clang__)
338
+ if (grx_detect_simd() >= 2) {
339
+ grx_add_scalar_avx2(a, s, out, n);
340
+ return;
341
+ }
160
342
  #endif
343
+ for (size_t i = 0; i < n; i++) out[i] = a[i] + s;
161
344
  }
162
345
 
163
- GRX_API void grx_negate(const double * restrict a, double * restrict out, size_t n) {
346
+ GRX_API void grx_negate(const double *a, double *out, size_t n) {
164
347
  grx_scale(a, -1.0, out, n);
165
348
  }
166
349
 
167
350
  /* ================================================================
168
- * MATEMÁTICAS ELEMENT-WISE
351
+ * MATEMATICAS ELEMENT-WISE
169
352
  * ================================================================ */
170
353
 
171
- GRX_API void grx_abs(const double * restrict a, double * restrict out, size_t n) {
172
- #if defined(GRX_AVX2_FMA) || defined(GRX_AVX2)
173
- /* Máscara para limpiar el bit de signo (AND con 0x7FFFFFFFFFFFFFFF) */
174
- __m256d mask = _mm256_castsi256_pd(
175
- _mm256_set1_epi64x(0x7FFFFFFFFFFFFFFFLL));
176
- size_t i = 0;
177
- for (; i + 4 <= n; i += 4)
178
- VST(out+i, _mm256_and_pd(VLD(a+i), mask));
179
- for (; i < n; i++) out[i] = fabs(a[i]);
180
- #else
354
+ GRX_API void grx_abs(const double *a, double *out, size_t n) {
181
355
  for (size_t i = 0; i < n; i++) out[i] = fabs(a[i]);
182
- #endif
183
356
  }
184
357
 
185
- GRX_API void grx_sqrt(const double * restrict a, double * restrict out, size_t n) {
186
- #if defined(GRX_AVX2_FMA) || defined(GRX_AVX2)
187
- size_t i = 0;
188
- for (; i + 4 <= n; i += 4)
189
- VST(out+i, _mm256_sqrt_pd(VLD(a+i)));
190
- for (; i < n; i++) out[i] = sqrt(a[i]);
191
- #else
358
+ GRX_API void grx_sqrt(const double *a, double *out, size_t n) {
192
359
  for (size_t i = 0; i < n; i++) out[i] = sqrt(a[i]);
193
- #endif
194
360
  }
195
361
 
196
- GRX_API void grx_square(const double * restrict a, double * restrict out, size_t n) {
362
+ GRX_API void grx_square(const double *a, double *out, size_t n) {
197
363
  grx_mul(a, a, out, n);
198
364
  }
199
365
 
200
- GRX_API void grx_log(const double * restrict a, double * restrict out, size_t n) {
201
- /* log no tiene intrínseco SIMD estándar; -ffast-math + -march=native
202
- * permite al compilador auto-vectorizar con SVML si está disponible */
366
+ GRX_API void grx_log(const double *a, double *out, size_t n) {
203
367
  for (size_t i = 0; i < n; i++) out[i] = log(a[i]);
204
368
  }
205
369
 
206
- GRX_API void grx_exp(const double * restrict a, double * restrict out, size_t n) {
370
+ GRX_API void grx_exp(const double *a, double *out, size_t n) {
207
371
  for (size_t i = 0; i < n; i++) out[i] = exp(a[i]);
208
372
  }
209
373
 
210
- GRX_API void grx_pow(const double * restrict a, double e,
211
- double * restrict out, size_t n) {
374
+ GRX_API void grx_pow(const double *a, double e, double *out, size_t n) {
212
375
  for (size_t i = 0; i < n; i++) out[i] = pow(a[i], e);
213
376
  }
214
377
 
215
- GRX_API void grx_clip(const double * restrict a, double lo, double hi,
216
- double * restrict out, size_t n) {
217
- #if defined(GRX_AVX2_FMA) || defined(GRX_AVX2)
218
- __m256d vlo = _mm256_set1_pd(lo);
219
- __m256d vhi = _mm256_set1_pd(hi);
220
- size_t i = 0;
221
- for (; i + 4 <= n; i += 4)
222
- VST(out+i, _mm256_min_pd(_mm256_max_pd(VLD(a+i), vlo), vhi));
223
- for (; i < n; i++) out[i] = a[i] < lo ? lo : (a[i] > hi ? hi : a[i]);
224
- #else
225
- for (size_t i = 0; i < n; i++)
226
- out[i] = a[i] < lo ? lo : (a[i] > hi ? hi : a[i]);
227
- #endif
378
+ GRX_API void grx_clip(const double *a, double lo, double hi, double *out, size_t n) {
379
+ for (size_t i = 0; i < n; i++) {
380
+ double v = a[i];
381
+ out[i] = (v < lo) ? lo : ((v > hi) ? hi : v);
382
+ }
228
383
  }
229
384
 
230
385
  /* ================================================================
231
386
  * REDUCCIONES
232
387
  * ================================================================ */
233
388
 
234
- GRX_API double grx_sum(const double * restrict a, size_t n) {
235
- double acc = 0.0;
236
- #ifdef GRX_AVX2_FMA
237
- __m256d v0 = _mm256_setzero_pd(), v1 = _mm256_setzero_pd();
238
- size_t i = 0;
239
- for (; i + 8 <= n; i += 8) {
240
- v0 = _mm256_add_pd(v0, VLD(a+i));
241
- v1 = _mm256_add_pd(v1, VLD(a+i+4));
389
+ GRX_API double grx_sum(const double *a, size_t n) {
390
+ #if defined(__GNUC__) || defined(__clang__)
391
+ if (grx_detect_simd() >= 2) {
392
+ return grx_sum_avx2(a, n);
242
393
  }
243
- v0 = _mm256_add_pd(v0, v1);
244
- for (; i + 4 <= n; i += 4) v0 = _mm256_add_pd(v0, VLD(a+i));
245
- double tmp[4]; _mm256_store_pd(tmp, v0);
246
- acc = tmp[0] + tmp[1] + tmp[2] + tmp[3];
247
- for (; i < n; i++) acc += a[i];
248
- #elif defined(GRX_AVX2)
249
- __m256d vacc = _mm256_setzero_pd();
250
- size_t i = 0;
251
- for (; i + 4 <= n; i += 4) vacc = _mm256_add_pd(vacc, VLD(a+i));
252
- double tmp[4]; _mm256_storeu_pd(tmp, vacc);
253
- acc = tmp[0] + tmp[1] + tmp[2] + tmp[3];
254
- for (; i < n; i++) acc += a[i];
255
- #else
256
- for (size_t i = 0; i < n; i++) acc += a[i];
257
394
  #endif
395
+ double acc = 0.0;
396
+ for (size_t i = 0; i < n; i++) acc += a[i];
258
397
  return acc;
259
398
  }
260
399
 
261
- GRX_API double grx_mean(const double * restrict a, size_t n) {
262
- return n > 0 ? grx_sum(a, n) / (double)n : 0.0;
400
+ GRX_API double grx_mean(const double *a, size_t n) {
401
+ if (n == 0) return 0.0;
402
+ return grx_sum(a, n) / (double)n;
263
403
  }
264
404
 
265
- GRX_API double grx_max(const double * restrict a, size_t n) {
266
- if (n == 0) return -DBL_MAX;
405
+ GRX_API double grx_max(const double *a, size_t n) {
406
+ if (n == 0) return 0.0;
267
407
  double m = a[0];
268
- for (size_t i = 1; i < n; i++) if (a[i] > m) m = a[i];
408
+ for (size_t i = 1; i < n; i++) {
409
+ if (a[i] > m) m = a[i];
410
+ }
269
411
  return m;
270
412
  }
271
413
 
272
- GRX_API double grx_min(const double * restrict a, size_t n) {
273
- if (n == 0) return DBL_MAX;
414
+ GRX_API double grx_min(const double *a, size_t n) {
415
+ if (n == 0) return 0.0;
274
416
  double m = a[0];
275
- for (size_t i = 1; i < n; i++) if (a[i] < m) m = a[i];
417
+ for (size_t i = 1; i < n; i++) {
418
+ if (a[i] < m) m = a[i];
419
+ }
276
420
  return m;
277
421
  }
278
422
 
279
423
  /* ================================================================
280
- * ÁLGEBRA LINEAL
424
+ * ALGEBRA LINEAL
281
425
  * ================================================================ */
282
426
 
283
- GRX_API double grx_dot(const double * restrict a, const double * restrict b, size_t n) {
284
- double acc = 0.0;
285
- #ifdef GRX_AVX2_FMA
286
- __m256d v0 = _mm256_setzero_pd(), v1 = _mm256_setzero_pd();
287
- size_t i = 0;
288
- for (; i + 8 <= n; i += 8) {
289
- v0 = _mm256_fmadd_pd(VLD(a+i), VLD(b+i), v0);
290
- v1 = _mm256_fmadd_pd(VLD(a+i+4), VLD(b+i+4), v1);
427
+ GRX_API double grx_dot(const double *a, const double *b, size_t n) {
428
+ #if defined(__GNUC__) || defined(__clang__)
429
+ if (grx_detect_simd() >= 2) {
430
+ return grx_dot_avx2(a, b, n);
291
431
  }
292
- v0 = _mm256_add_pd(v0, v1);
293
- for (; i + 4 <= n; i += 4) v0 = _mm256_fmadd_pd(VLD(a+i), VLD(b+i), v0);
294
- double tmp[4]; _mm256_store_pd(tmp, v0);
295
- acc = tmp[0] + tmp[1] + tmp[2] + tmp[3];
296
- for (; i < n; i++) acc += a[i] * b[i];
297
- #elif defined(GRX_AVX2)
298
- __m256d vacc = _mm256_setzero_pd();
299
- size_t i = 0;
300
- for (; i + 4 <= n; i += 4)
301
- vacc = _mm256_add_pd(vacc, _mm256_mul_pd(VLD(a+i), VLD(b+i)));
302
- double tmp[4]; _mm256_storeu_pd(tmp, vacc);
303
- acc = tmp[0] + tmp[1] + tmp[2] + tmp[3];
304
- for (; i < n; i++) acc += a[i] * b[i];
305
- #else
306
- for (size_t i = 0; i < n; i++) acc += a[i] * b[i];
307
432
  #endif
433
+ double acc = 0.0;
434
+ for (size_t i = 0; i < n; i++) acc += a[i] * b[i];
308
435
  return acc;
309
436
  }
310
437
 
311
- /* matmul con tiling cache-friendly */
312
- GRX_API void grx_matmul(const double * restrict a, const double * restrict b,
313
- double * restrict out, size_t M, size_t K, size_t N) {
438
+ GRX_API void grx_matmul(const double *a, const double *b, double *out,
439
+ size_t M, size_t K, size_t N) {
314
440
  memset(out, 0, M * N * sizeof(double));
315
- for (size_t ii = 0; ii < M; ii += TILE) {
316
- size_t ie = ii + TILE < M ? ii + TILE : M;
317
- for (size_t kk = 0; kk < K; kk += TILE) {
318
- size_t ke = kk + TILE < K ? kk + TILE : K;
319
- for (size_t jj = 0; jj < N; jj += TILE) {
320
- size_t je = jj + TILE < N ? jj + TILE : N;
321
- for (size_t i = ii; i < ie; i++)
322
- for (size_t k = kk; k < ke; k++) {
323
- double aik = a[i*K+k];
324
- for (size_t j = jj; j < je; j++)
325
- out[i*N+j] += aik * b[k*N+j];
441
+
442
+ /* Multiplicacion de matrices optimizada por bloques con cache tiling */
443
+ for (size_t bi = 0; bi < M; bi += TILE) {
444
+ size_t imax = bi + TILE < M ? bi + TILE : M;
445
+ for (size_t bk = 0; bk < K; bk += TILE) {
446
+ size_t kmax = bk + TILE < K ? bk + TILE : K;
447
+ for (size_t bj = 0; bj < N; bj += TILE) {
448
+ size_t jmax = bj + TILE < N ? bj + TILE : N;
449
+
450
+ for (size_t i = bi; i < imax; i++) {
451
+ for (size_t k = bk; k < kmax; k++) {
452
+ double a_ik = a[i * K + k];
453
+ double *out_i = out + i * N;
454
+ const double *b_k = b + k * N;
455
+ for (size_t j = bj; j < jmax; j++) {
456
+ out_i[j] += a_ik * b_k[j];
457
+ }
326
458
  }
459
+ }
327
460
  }
328
461
  }
329
462
  }
@@ -333,166 +466,111 @@ GRX_API void grx_matmul(const double * restrict a, const double * restrict b,
333
466
  * ACTIVACIONES
334
467
  * ================================================================ */
335
468
 
336
- GRX_API void grx_relu(const double * restrict a, double * restrict out, size_t n) {
337
- #if defined(GRX_AVX2_FMA) || defined(GRX_AVX2)
338
- __m256d vz = _mm256_setzero_pd();
339
- size_t i = 0;
340
- for (; i + 8 <= n; i += 8) {
341
- VST(out+i, _mm256_max_pd(VLD(a+i), vz));
342
- VST(out+i+4, _mm256_max_pd(VLD(a+i+4), vz));
469
+ GRX_API void grx_relu(const double *a, double *out, size_t n) {
470
+ #if defined(__GNUC__) || defined(__clang__)
471
+ if (grx_detect_simd() >= 2) {
472
+ grx_relu_avx2(a, out, n);
473
+ return;
343
474
  }
344
- for (; i + 4 <= n; i += 4) VST(out+i, _mm256_max_pd(VLD(a+i), vz));
345
- for (; i < n; i++) out[i] = a[i] > 0.0 ? a[i] : 0.0;
346
- #else
347
- for (size_t i = 0; i < n; i++) out[i] = a[i] > 0.0 ? a[i] : 0.0;
348
475
  #endif
476
+ for (size_t i = 0; i < n; i++) out[i] = a[i] > 0.0 ? a[i] : 0.0;
349
477
  }
350
478
 
351
- GRX_API void grx_leaky_relu(const double * restrict a, double alpha,
352
- double * restrict out, size_t n) {
353
- #if defined(GRX_AVX2_FMA) || defined(GRX_AVX2)
354
- __m256d va = _mm256_set1_pd(alpha);
355
- __m256d vz = _mm256_setzero_pd();
356
- size_t i = 0;
357
- for (; i + 4 <= n; i += 4) {
358
- __m256d v = VLD(a+i);
359
- /* max(v, alpha*v): si v>0 → v, si v<=0 → alpha*v */
360
- VST(out+i, _mm256_blendv_pd(_mm256_mul_pd(v, va), v,
361
- _mm256_cmp_pd(v, vz, _CMP_GT_OQ)));
479
+ GRX_API void grx_leaky_relu(const double *a, double alpha, double *out, size_t n) {
480
+ for (size_t i = 0; i < n; i++) {
481
+ double v = a[i];
482
+ out[i] = v >= 0.0 ? v : alpha * v;
362
483
  }
363
- for (; i < n; i++) out[i] = a[i] > 0.0 ? a[i] : alpha * a[i];
364
- #else
365
- for (size_t i = 0; i < n; i++) out[i] = a[i] > 0.0 ? a[i] : alpha * a[i];
366
- #endif
367
484
  }
368
485
 
369
- GRX_API void grx_tanh_act(const double * restrict a, double * restrict out, size_t n) {
486
+ GRX_API void grx_tanh_act(const double *a, double *out, size_t n) {
370
487
  for (size_t i = 0; i < n; i++) out[i] = tanh(a[i]);
371
488
  }
372
489
 
373
- GRX_API void grx_sigmoid(const double * restrict a, double * restrict out, size_t n) {
374
- for (size_t i = 0; i < n; i++) out[i] = 1.0 / (1.0 + exp(-a[i]));
490
+ GRX_API void grx_sigmoid(const double *a, double *out, size_t n) {
491
+ for (size_t i = 0; i < n; i++) {
492
+ double v = a[i];
493
+ if (v >= 0.0) {
494
+ double ev = exp(-v);
495
+ out[i] = 1.0 / (1.0 + ev);
496
+ } else {
497
+ double ev = exp(v);
498
+ out[i] = ev / (1.0 + ev);
499
+ }
500
+ }
375
501
  }
376
502
 
377
- GRX_API void grx_softmax(const double * restrict a, double * restrict out, size_t n) {
378
- double max_val = grx_max(a, n);
503
+ GRX_API void grx_softmax(const double *a, double *out, size_t n) {
504
+ if (n == 0) return;
505
+ double max_v = a[0];
506
+ for (size_t i = 1; i < n; i++) {
507
+ if (a[i] > max_v) max_v = a[i];
508
+ }
379
509
  double sum = 0.0;
380
- for (size_t i = 0; i < n; i++) { out[i] = exp(a[i] - max_val); sum += out[i]; }
510
+ for (size_t i = 0; i < n; i++) {
511
+ double ev = exp(a[i] - max_v);
512
+ out[i] = ev;
513
+ sum += ev;
514
+ }
381
515
  double inv = 1.0 / sum;
382
516
  for (size_t i = 0; i < n; i++) out[i] *= inv;
383
517
  }
384
518
 
385
519
  /* ================================================================
386
- * OPTIMIZADORES (in-place sobre parámetros)
520
+ * OPTIMIZADORES IN-PLACE
387
521
  * ================================================================ */
388
522
 
389
- /* SGD: param[i] -= lr * grad[i] */
390
- GRX_API void grx_sgd_step(double * restrict param, const double * restrict grad,
391
- double lr, size_t n) {
392
- #ifdef GRX_AVX2_FMA
393
- __m256d vlr = _mm256_set1_pd(lr);
394
- size_t i = 0;
395
- for (; i + 8 <= n; i += 8) {
396
- /* param -= lr * grad usando FMA: param = -lr*grad + param */
397
- VST(param+i, _mm256_fnmadd_pd(vlr, VLD(grad+i), VLD(param+i)));
398
- VST(param+i+4, _mm256_fnmadd_pd(vlr, VLD(grad+i+4), VLD(param+i+4)));
523
+ GRX_API void grx_sgd_step(double *param, const double *grad,
524
+ double lr, size_t n) {
525
+ for (size_t i = 0; i < n; i++) {
526
+ param[i] -= lr * grad[i];
399
527
  }
400
- for (; i + 4 <= n; i += 4)
401
- VST(param+i, _mm256_fnmadd_pd(vlr, VLD(grad+i), VLD(param+i)));
402
- for (; i < n; i++) param[i] -= lr * grad[i];
403
- #else
404
- for (size_t i = 0; i < n; i++) param[i] -= lr * grad[i];
405
- #endif
406
528
  }
407
529
 
408
- /*
409
- * Adam: Kingma & Ba 2015
410
- * m = beta1*m + (1-beta1)*grad
411
- * v = beta2*v + (1-beta2)*grad^2
412
- * m_hat = m / (1 - beta1^t)
413
- * v_hat = v / (1 - beta2^t)
414
- * param -= lr * m_hat / (sqrt(v_hat) + eps)
415
- *
416
- * beta1t = beta1^t (pasado desde Ruby, se actualiza por paso)
417
- * beta2t = beta2^t
418
- */
419
- GRX_API void grx_adam_step(double * restrict param,
420
- double * restrict m, double * restrict v,
421
- const double * restrict grad,
422
- double lr, double beta1, double beta2,
423
- double epsilon, double beta1t, double beta2t,
424
- size_t n) {
425
- double one_m_b1 = 1.0 - beta1;
426
- double one_m_b2 = 1.0 - beta2;
427
- double inv_1mb1t = 1.0 / (1.0 - beta1t);
428
- double inv_1mb2t = 1.0 / (1.0 - beta2t);
429
-
430
- #ifdef GRX_AVX2_FMA
431
- __m256d vb1 = _mm256_set1_pd(beta1);
432
- __m256d vb2 = _mm256_set1_pd(beta2);
433
- __m256d v1mb1 = _mm256_set1_pd(one_m_b1);
434
- __m256d v1mb2 = _mm256_set1_pd(one_m_b2);
435
- __m256d vlr = _mm256_set1_pd(lr);
436
- __m256d veps = _mm256_set1_pd(epsilon);
437
- __m256d vi1b1t = _mm256_set1_pd(inv_1mb1t);
438
- __m256d vi2b2t = _mm256_set1_pd(inv_1mb2t);
439
-
440
- size_t i = 0;
441
- for (; i + 4 <= n; i += 4) {
442
- __m256d g = VLD(grad+i);
443
- /* m = beta1*m + (1-beta1)*g */
444
- __m256d mi = _mm256_fmadd_pd(vb1, VLD(m+i), _mm256_mul_pd(v1mb1, g));
445
- /* v = beta2*v + (1-beta2)*g^2 */
446
- __m256d vi = _mm256_fmadd_pd(vb2, VLD(v+i),
447
- _mm256_mul_pd(v1mb2, _mm256_mul_pd(g, g)));
448
- VST(m+i, mi);
449
- VST(v+i, vi);
450
- /* m_hat = m / (1-beta1^t), v_hat = v / (1-beta2^t) */
451
- __m256d mh = _mm256_mul_pd(mi, vi1b1t);
452
- __m256d vh = _mm256_mul_pd(vi, vi2b2t);
453
- /* param -= lr * mh / (sqrt(vh) + eps) */
454
- __m256d denom = _mm256_add_pd(_mm256_sqrt_pd(vh), veps);
455
- VST(param+i, _mm256_fnmadd_pd(vlr, _mm256_div_pd(mh, denom), VLD(param+i)));
456
- }
457
- for (; i < n; i++) {
458
- m[i] = beta1 * m[i] + one_m_b1 * grad[i];
459
- v[i] = beta2 * v[i] + one_m_b2 * grad[i] * grad[i];
460
- double mh = m[i] * inv_1mb1t;
461
- double vh = v[i] * inv_1mb2t;
462
- param[i] -= lr * mh / (sqrt(vh) + epsilon);
530
+ GRX_API void grx_adam_step(double *param, double *m, double *v,
531
+ const double *grad, double lr,
532
+ double beta1, double beta2, double epsilon,
533
+ double beta1t, double beta2t, size_t n) {
534
+ #if defined(__GNUC__) || defined(__clang__)
535
+ if (grx_detect_simd() >= 2) {
536
+ grx_adam_step_avx2(param, m, v, grad, lr, beta1, beta2, epsilon, beta1t, beta2t, n);
537
+ return;
463
538
  }
464
- #else
539
+ #endif
540
+ double cb1 = 1.0 / (1.0 - beta1t);
541
+ double cb2 = 1.0 / (1.0 - beta2t);
465
542
  for (size_t i = 0; i < n; i++) {
466
- m[i] = beta1 * m[i] + one_m_b1 * grad[i];
467
- v[i] = beta2 * v[i] + one_m_b2 * grad[i] * grad[i];
468
- double mh = m[i] * inv_1mb1t;
469
- double vh = v[i] * inv_1mb2t;
470
- param[i] -= lr * mh / (sqrt(vh) + epsilon);
543
+ double g = grad[i];
544
+ m[i] = beta1 * m[i] + (1.0 - beta1) * g;
545
+ v[i] = beta2 * v[i] + (1.0 - beta2) * g * g;
546
+ double m_hat = m[i] * cb1;
547
+ double v_hat = v[i] * cb2;
548
+ param[i] -= lr * m_hat / (sqrt(v_hat) + epsilon);
471
549
  }
472
- #endif
473
550
  }
474
551
 
475
552
  /* ================================================================
476
- * INICIALIZACIÓN DE PESOS
553
+ * INICIALIZACION DE PESOS
477
554
  * ================================================================ */
478
555
 
479
- /* LCG simple (no criptográfico, pero rápido y sin dependencias) */
480
556
  static uint64_t grx_rng_state = 0;
481
557
 
482
558
  static void grx_rng_seed(void) {
483
- grx_rng_state = (uint64_t)time(NULL) ^ (uint64_t)(uintptr_t)&grx_rng_state;
559
+ if (grx_rng_state == 0) {
560
+ grx_rng_state = (uint64_t)time(NULL) ^ (uint64_t)(uintptr_t)&grx_rng_state ^ 0x853c49e6748fea9bULL;
561
+ if (grx_rng_state == 0) grx_rng_state = 1;
562
+ }
484
563
  }
485
564
 
486
- /* Genera double uniforme en [0, 1) */
565
+ /* Generador xorshift64* puramente escalar y universal */
487
566
  static double grx_rand01(void) {
488
- /* xorshift64 */
489
567
  grx_rng_state ^= grx_rng_state << 13;
490
568
  grx_rng_state ^= grx_rng_state >> 7;
491
569
  grx_rng_state ^= grx_rng_state << 17;
492
- return (double)(grx_rng_state >> 11) / (double)(1ULL << 53);
570
+ uint64_t val = grx_rng_state * 0x2545F4914F6CDD1DULL;
571
+ return (double)(val >> 11) * (1.0 / 9007199254740992.0); /* 2^53 */
493
572
  }
494
573
 
495
- /* Box-Muller: genera par de normales N(0,1) */
496
574
  static void grx_box_muller(double *z0, double *z1) {
497
575
  double u1, u2;
498
576
  do { u1 = grx_rand01(); } while (u1 < 1e-15);
@@ -503,16 +581,17 @@ static void grx_box_muller(double *z0, double *z1) {
503
581
  }
504
582
 
505
583
  GRX_API void grx_init_xavier_uniform(double *out, size_t n,
506
- size_t fan_in, size_t fan_out) {
584
+ size_t fan_in, size_t fan_out) {
507
585
  grx_rng_seed();
508
586
  double limit = sqrt(6.0 / (double)(fan_in + fan_out));
509
- for (size_t i = 0; i < n; i++)
587
+ for (size_t i = 0; i < n; i++) {
510
588
  out[i] = (grx_rand01() * 2.0 - 1.0) * limit;
589
+ }
511
590
  }
512
591
 
513
592
  GRX_API void grx_init_he_normal(double *out, size_t n, size_t fan_in) {
514
593
  grx_rng_seed();
515
- double std = sqrt(2.0 / (double)fan_in);
594
+ double std = sqrt(2.0 / (double)(fan_in > 0 ? fan_in : 1));
516
595
  size_t i = 0;
517
596
  for (; i + 1 < n; i += 2) {
518
597
  double z0, z1;
@@ -527,8 +606,9 @@ GRX_API void grx_init_he_normal(double *out, size_t n, size_t fan_in) {
527
606
  }
528
607
  }
529
608
 
530
- /* ============================================================
531
- * RUBY EXTENSION INIT — requerido por rake-compiler / mkmf
532
- * No hace nada: la librería se carga vía Fiddle, no como extensión Ruby nativa.
533
- * ============================================================ */
534
- void Init_grx_core(void) { /* no-op */ }
609
+ /* ================================================================
610
+ * RUBY EXTENSION ENTRYPOINT
611
+ * ================================================================ */
612
+ GRX_API void Init_grx_core(void) {
613
+ grx_detect_simd();
614
+ }