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.
- checksums.yaml +4 -4
- data/CHANGELOG.md +9 -0
- data/GUIA_PRINCIPIANTES.md +307 -24
- data/README.es.md +1090 -161
- data/README.md +1102 -162
- data/ext/grx/extconf.rb +4 -18
- data/ext/grx/grx_core.c +411 -331
- data/ext/grx/grx_core.h +23 -13
- data/ext/unix/Makefile +3 -27
- data/ext/windows/Makefile.mingw +3 -23
- data/grx-tensor.gemspec +37 -33
- data/lib/grx/c_api.rb +47 -15
- data/lib/grx/nn.rb +9 -7
- data/lib/grx/optim.rb +12 -7
- data/lib/grx/storage.rb +6 -5
- data/lib/grx/tensor.rb +36 -0
- data/lib/grx/utils.rb +15 -0
- data/lib/grx/version.rb +1 -1
- data/lib/grx.rb +8 -3
- metadata +31 -27
data/ext/grx/grx_core.c
CHANGED
|
@@ -1,18 +1,20 @@
|
|
|
1
1
|
/*
|
|
2
|
-
* grx_core.c —
|
|
3
|
-
*
|
|
4
|
-
*
|
|
5
|
-
* -
|
|
6
|
-
* -
|
|
7
|
-
*
|
|
8
|
-
*
|
|
9
|
-
* -
|
|
10
|
-
* -
|
|
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
|
|
15
|
-
#define _POSIX_C_SOURCE 200809L
|
|
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(
|
|
30
|
+
#if defined(__GNUC__) || defined(__clang__)
|
|
31
|
+
#define GRX_TARGET_AVX2 __attribute__((target("avx2,fma")))
|
|
29
32
|
#include <immintrin.h>
|
|
30
|
-
|
|
31
|
-
#
|
|
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
|
-
*
|
|
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)
|
|
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
|
-
|
|
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
|
-
*
|
|
96
|
+
* ARITMETICA ELEMENT-WISE: KERNELS AVX2 Y ESCALARES
|
|
66
97
|
* ================================================================ */
|
|
67
98
|
|
|
68
|
-
|
|
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
|
-
|
|
82
|
-
|
|
83
|
-
|
|
84
|
-
|
|
85
|
-
|
|
86
|
-
|
|
87
|
-
|
|
88
|
-
for (; i
|
|
89
|
-
|
|
90
|
-
|
|
91
|
-
|
|
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
|
-
|
|
105
|
-
|
|
106
|
-
|
|
107
|
-
|
|
108
|
-
|
|
109
|
-
|
|
110
|
-
|
|
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
|
-
|
|
114
|
-
|
|
115
|
-
|
|
116
|
-
|
|
117
|
-
|
|
118
|
-
|
|
119
|
-
|
|
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
|
-
|
|
123
|
-
|
|
124
|
-
|
|
125
|
-
|
|
126
|
-
|
|
127
|
-
|
|
128
|
-
|
|
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
|
-
|
|
132
|
-
|
|
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
|
-
|
|
138
|
-
|
|
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
|
-
|
|
148
|
-
|
|
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
|
-
|
|
154
|
-
|
|
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
|
-
|
|
159
|
-
|
|
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 *
|
|
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
|
-
*
|
|
351
|
+
* MATEMATICAS ELEMENT-WISE
|
|
169
352
|
* ================================================================ */
|
|
170
353
|
|
|
171
|
-
GRX_API void grx_abs(const double *
|
|
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 *
|
|
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 *
|
|
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 *
|
|
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 *
|
|
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 *
|
|
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 *
|
|
216
|
-
|
|
217
|
-
|
|
218
|
-
|
|
219
|
-
|
|
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 *
|
|
235
|
-
|
|
236
|
-
|
|
237
|
-
|
|
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 *
|
|
262
|
-
|
|
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 *
|
|
266
|
-
if (n == 0) return
|
|
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++)
|
|
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 *
|
|
273
|
-
if (n == 0) return
|
|
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++)
|
|
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
|
-
*
|
|
424
|
+
* ALGEBRA LINEAL
|
|
281
425
|
* ================================================================ */
|
|
282
426
|
|
|
283
|
-
GRX_API double grx_dot(const double *
|
|
284
|
-
|
|
285
|
-
|
|
286
|
-
|
|
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
|
-
|
|
312
|
-
|
|
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
|
-
|
|
316
|
-
|
|
317
|
-
|
|
318
|
-
|
|
319
|
-
|
|
320
|
-
|
|
321
|
-
|
|
322
|
-
|
|
323
|
-
|
|
324
|
-
|
|
325
|
-
|
|
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 *
|
|
337
|
-
#if defined(
|
|
338
|
-
|
|
339
|
-
|
|
340
|
-
|
|
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 *
|
|
352
|
-
|
|
353
|
-
|
|
354
|
-
|
|
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 *
|
|
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 *
|
|
374
|
-
for (size_t i = 0; i < n; 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 *
|
|
378
|
-
|
|
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++) {
|
|
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
|
|
520
|
+
* OPTIMIZADORES IN-PLACE
|
|
387
521
|
* ================================================================ */
|
|
388
522
|
|
|
389
|
-
|
|
390
|
-
|
|
391
|
-
|
|
392
|
-
|
|
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
|
-
*
|
|
410
|
-
|
|
411
|
-
|
|
412
|
-
|
|
413
|
-
|
|
414
|
-
|
|
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
|
-
#
|
|
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
|
-
|
|
467
|
-
|
|
468
|
-
|
|
469
|
-
double
|
|
470
|
-
|
|
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
|
-
*
|
|
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
|
|
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
|
-
/*
|
|
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
|
-
|
|
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
|
-
|
|
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
|
|
532
|
-
*
|
|
533
|
-
|
|
534
|
-
|
|
609
|
+
/* ================================================================
|
|
610
|
+
* RUBY EXTENSION ENTRYPOINT
|
|
611
|
+
* ================================================================ */
|
|
612
|
+
GRX_API void Init_grx_core(void) {
|
|
613
|
+
grx_detect_simd();
|
|
614
|
+
}
|