yum-archive/TaSTT-Whisper

High-performance GPGPU inference of OpenAI's Whisper automatic speech recognition (ASR) model

git clone https://git.yummers.dev/yum-archive/TaSTT-Whisper

KonstantinSource codes8c4603c

master
244.3 KiB8336 linesraw
1#include "ggml.h"
2
3#if defined(_MSC_VER) || defined(__MINGW32__)
4#include <malloc.h> // using malloc.h with MSC/MINGW
5#elif !defined(__FreeBSD__)
6#include <alloca.h>
7#endif
8
9#include <assert.h>
10#include <time.h>
11#include <math.h>
12#include <stdlib.h>
13#include <string.h>
14#include <stdint.h>
15#include <stdio.h>
16
17// if C99 - static_assert is noop
18// ref: https://stackoverflow.com/a/53923785/4039976
19#ifndef static_assert
20#define static_assert(cond, msg) struct global_scope_noop_trick
21#endif
22
23#if defined _MSC_VER || defined(__MINGW32__)
24
25#if !defined(__MINGW32__)
26#include <Windows.h>
27#else
28// ref: https://github.com/ggerganov/whisper.cpp/issues/168
29#include <windows.h>
30#include <errno.h>
31#endif
32
33typedef volatile LONG atomic_int;
34typedef atomic_int atomic_bool;
35
36static void atomic_store(atomic_int* ptr, LONG val) {
37    InterlockedExchange(ptr, val);
38}
39static LONG atomic_load(atomic_int* ptr) {
40    return InterlockedCompareExchange(ptr, 0, 0);
41}
42static LONG atomic_fetch_add(atomic_int* ptr, LONG inc) {
43    return InterlockedExchangeAdd(ptr, inc);
44}
45static LONG atomic_fetch_sub(atomic_int* ptr, LONG dec) {
46    return atomic_fetch_add(ptr, -(dec));
47}
48
49typedef HANDLE pthread_t;
50
51typedef DWORD thread_ret_t;
52static int pthread_create(pthread_t* out, void* unused, thread_ret_t(*func)(void*), void* arg) {
53    HANDLE handle = CreateThread(NULL, 0, (LPTHREAD_START_ROUTINE) func, arg, 0, NULL);
54    if (handle == NULL)
55    {
56        return EAGAIN;
57    }
58
59    *out = handle;
60    return 0;
61}
62
63static int pthread_join(pthread_t thread, void* unused) {
64    return (int) WaitForSingleObject(thread, INFINITE);
65}
66
67static int sched_yield (void) {
68    Sleep (0);
69    return 0;
70}
71#else
72#include <pthread.h>
73#include <stdatomic.h>
74
75typedef void* thread_ret_t;
76#endif
77
78#ifdef __HAIKU__
79#define static_assert(cond, msg) _Static_assert(cond, msg)
80#endif
81
82#define GGML_DEBUG 0
83#define GGML_GELU_FP16
84
85#if UINTPTR_MAX == 0xFFFFFFFF
86    #define GGML_MEM_ALIGN 4
87#else
88    #define GGML_MEM_ALIGN 16
89#endif
90
91#define MAX(a, b) ((a) > (b) ? (a) : (b))
92#define MIN(a, b) ((a) < (b) ? (a) : (b))
93
94#define UNUSED(x) (void)(x)
95#define SWAP(x, y, T) do { T SWAP = x; x = y; y = SWAP; } while (0)
96
97#define GGML_ASSERT(x) \
98    do { \
99        if (!(x)) { \
100            logError( u8"GGML_ASSERT: %s:%d: %s", __FILE__, __LINE__, #x); \
101            abort(); \
102        } \
103    } while (0)
104
105#ifdef GGML_USE_ACCELERATE
106#include <Accelerate/Accelerate.h>
107#elif GGML_USE_OPENBLAS
108#include <cblas.h>
109#endif
110
111// floating point type used to accumulate sums
112typedef double ggml_float;
113
114// 16-bit float
115// on Arm, we use __fp16
116// on x86, we use uint16_t
117#ifdef __ARM_NEON
118
119// if YCM cannot find <arm_neon.h>, make a symbolic link to it, for example:
120//
121//   $ ln -sfn /Library/Developer/CommandLineTools/usr/lib/clang/13.1.6/include/arm_neon.h ./src/
122//
123#include <arm_neon.h>
124
125float ggml_fp16_to_fp32(ggml_fp16_t x) {
126    return x;
127}
128
129ggml_fp16_t ggml_fp32_to_fp16(float x) {
130    return x;
131}
132
133#define GGML_FP16_TO_FP32(x) (x)
134#define GGML_FP32_TO_FP16(x) (x)
135
136#else
137
138#ifdef __wasm_simd128__
139#include <wasm_simd128.h>
140#else
141#ifdef __POWER9_VECTOR__
142#include <altivec.h>
143#undef bool
144#define bool _Bool
145#else
146#include <immintrin.h>
147#endif
148#endif
149
150#ifdef __F16C__
151float ggml_fp16_to_fp32(ggml_fp16_t h) {
152    return _cvtsh_ss(h);
153}
154ggml_fp16_t ggml_fp32_to_fp16(float f) {
155    return _cvtss_sh(f, 0);
156}
157
158#define GGML_FP16_TO_FP32(x) _cvtsh_ss(x)
159#define GGML_FP32_TO_FP16(x) _cvtss_sh(x, 0)
160
161#else
162
163// FP16 <-> FP32
164// ref: https://github.com/Maratyszcza/FP16
165
166static inline float fp32_from_bits(uint32_t w) {
167    union {
168        uint32_t as_bits;
169        float as_value;
170    } fp32;
171    fp32.as_bits = w;
172    return fp32.as_value;
173}
174
175static inline uint32_t fp32_to_bits(float f) {
176	union {
177		float as_value;
178		uint32_t as_bits;
179	} fp32;
180	fp32.as_value = f;
181	return fp32.as_bits;
182}
183
184float ggml_fp16_to_fp32(ggml_fp16_t h) {
185    const uint32_t w = (uint32_t) h << 16;
186    const uint32_t sign = w & UINT32_C(0x80000000);
187    const uint32_t two_w = w + w;
188
189    const uint32_t exp_offset = UINT32_C(0xE0) << 23;
190#if defined(__STDC_VERSION__) && (__STDC_VERSION__ >= 199901L) || defined(__GNUC__) && !defined(__STRICT_ANSI__)
191    const float exp_scale = 0x1.0p-112f;
192#else
193    const float exp_scale = fp32_from_bits(UINT32_C(0x7800000));
194#endif
195    const float normalized_value = fp32_from_bits((two_w >> 4) + exp_offset) * exp_scale;
196
197    const uint32_t magic_mask = UINT32_C(126) << 23;
198    const float magic_bias = 0.5f;
199    const float denormalized_value = fp32_from_bits((two_w >> 17) | magic_mask) - magic_bias;
200
201    const uint32_t denormalized_cutoff = UINT32_C(1) << 27;
202    const uint32_t result = sign |
203        (two_w < denormalized_cutoff ? fp32_to_bits(denormalized_value) : fp32_to_bits(normalized_value));
204    return fp32_from_bits(result);
205}
206
207ggml_fp16_t ggml_fp32_to_fp16(float f) {
208#if defined(__STDC_VERSION__) && (__STDC_VERSION__ >= 199901L) || defined(__GNUC__) && !defined(__STRICT_ANSI__)
209    const float scale_to_inf = 0x1.0p+112f;
210    const float scale_to_zero = 0x1.0p-110f;
211#else
212    const float scale_to_inf = fp32_from_bits(UINT32_C(0x77800000));
213    const float scale_to_zero = fp32_from_bits(UINT32_C(0x08800000));
214#endif
215    float base = (fabsf(f) * scale_to_inf) * scale_to_zero;
216
217    const uint32_t w = fp32_to_bits(f);
218    const uint32_t shl1_w = w + w;
219    const uint32_t sign = w & UINT32_C(0x80000000);
220    uint32_t bias = shl1_w & UINT32_C(0xFF000000);
221    if (bias < UINT32_C(0x71000000)) {
222        bias = UINT32_C(0x71000000);
223    }
224
225    base = fp32_from_bits((bias >> 1) + UINT32_C(0x07800000)) + base;
226    const uint32_t bits = fp32_to_bits(base);
227    const uint32_t exp_bits = (bits >> 13) & UINT32_C(0x00007C00);
228    const uint32_t mantissa_bits = bits & UINT32_C(0x00000FFF);
229    const uint32_t nonsign = exp_bits + mantissa_bits;
230    return (sign >> 16) | (shl1_w > UINT32_C(0xFF000000) ? UINT16_C(0x7E00) : nonsign);
231}
232
233#define GGML_FP16_TO_FP32(x) ggml_fp16_to_fp32(x)
234#define GGML_FP32_TO_FP16(x) ggml_fp32_to_fp16(x)
235
236#endif // __F16C__
237
238#endif // __ARM_NEON
239
240//
241// global data
242//
243
244// precomputed gelu table for f16 (128 KB)
245static ggml_fp16_t table_gelu_f16[1 << 16];
246
247// precomputed exp table for f16 (128 KB)
248static ggml_fp16_t table_exp_f16[1 << 16];
249
250//
251// timing
252//
253
254#if defined(_MSC_VER) || defined(__MINGW32__)
255static int64_t timer_freq;
256void ggml_time_init(void) {
257    LARGE_INTEGER frequency;
258    QueryPerformanceFrequency(&frequency);
259    timer_freq = frequency.QuadPart;
260}
261int64_t ggml_time_ms(void) {
262    LARGE_INTEGER t;
263    QueryPerformanceCounter(&t);
264    return (t.QuadPart * 1000) / timer_freq;
265}
266int64_t ggml_time_us(void) {
267    LARGE_INTEGER t;
268    QueryPerformanceCounter(&t);
269    return (t.QuadPart * 1000000) / timer_freq;
270}
271#else
272void ggml_time_init(void) {}
273int64_t ggml_time_ms(void) {
274    struct timespec ts;
275    clock_gettime(CLOCK_MONOTONIC, &ts);
276    return (int64_t)ts.tv_sec*1000 + (int64_t)ts.tv_nsec/1000000;
277}
278
279int64_t ggml_time_us(void) {
280    struct timespec ts;
281    clock_gettime(CLOCK_MONOTONIC, &ts);
282    return (int64_t)ts.tv_sec*1000000 + (int64_t)ts.tv_nsec/1000;
283}
284#endif
285
286int64_t ggml_cycles(void) {
287    return clock();
288}
289
290int64_t ggml_cycles_per_ms(void) {
291    return CLOCKS_PER_SEC/1000;
292}
293
294#ifdef GGML_PERF
295#define ggml_perf_time_ms()       ggml_time_ms()
296#define ggml_perf_time_us()       ggml_time_us()
297#define ggml_perf_cycles()        ggml_cycles()
298#define ggml_perf_cycles_per_ms() ggml_cycles_per_ms()
299#else
300#define ggml_perf_time_ms()       0
301#define ggml_perf_time_us()       0
302#define ggml_perf_cycles()        0
303#define ggml_perf_cycles_per_ms() 0
304#endif
305
306//
307// cache line
308//
309
310#if defined(__cpp_lib_hardware_interference_size)
311#define CACHE_LINE_SIZE hardware_destructive_interference_size
312#else
313#define CACHE_LINE_SIZE 64
314#endif
315
316static const size_t CACHE_LINE_SIZE_F32 = CACHE_LINE_SIZE/sizeof(float);
317
318//
319// simd mappings
320//
321
322// we define a common set of C macros which map to specific intrinsics based on the current architecture
323// we then implement the fundamental computation operations below using only these macros
324// adding support for new architectures requires to define the corresponding SIMD macros
325//
326// GGML_F32_STEP / GGML_F16_STEP
327//   number of elements to process in a single step
328//
329// GGML_F32_EPR / GGML_F16_EPR
330//   number of elements to fit in a single register
331//
332
333#if defined(__ARM_NEON) && defined(__ARM_FEATURE_FMA)
334
335#define GGML_SIMD
336
337// F32 NEON
338
339#define GGML_F32_STEP 16
340#define GGML_F32_EPR  4
341
342#define GGML_F32x4              float32x4_t
343#define GGML_F32x4_ZERO         vdupq_n_f32(0.0f)
344#define GGML_F32x4_SET1(x)      vdupq_n_f32(x)
345#define GGML_F32x4_LOAD         vld1q_f32
346#define GGML_F32x4_STORE        vst1q_f32
347#define GGML_F32x4_FMA(a, b, c) vfmaq_f32(a, b, c)
348#define GGML_F32x4_ADD          vaddq_f32
349#define GGML_F32x4_MUL          vmulq_f32
350#if defined(__ARM_FEATURE_QRDMX)
351    #define GGML_F32x4_REDUCE_ONE(x) vaddvq_f32(x)
352#else
353    #define GGML_F32x4_REDUCE_ONE(x) \
354    (vgetq_lane_f32(x, 0) +          \
355     vgetq_lane_f32(x, 1) +          \
356     vgetq_lane_f32(x, 2) +          \
357     vgetq_lane_f32(x, 3))
358#endif
359#define GGML_F32x4_REDUCE(res, x)              \
360{                                              \
361    for (int i = 0; i < GGML_F32_ARR/2; ++i) { \
362        x[2*i] = vaddq_f32(x[2*i], x[2*i+1]);  \
363    }                                          \
364    for (int i = 0; i < GGML_F32_ARR/4; ++i) { \
365        x[4*i] = vaddq_f32(x[4*i], x[4*i+2]);  \
366    }                                          \
367    for (int i = 0; i < GGML_F32_ARR/8; ++i) { \
368        x[8*i] = vaddq_f32(x[8*i], x[8*i+4]);  \
369    }                                          \
370    res = GGML_F32x4_REDUCE_ONE(x[0]);         \
371}
372
373#define GGML_F32_VEC        GGML_F32x4
374#define GGML_F32_VEC_ZERO   GGML_F32x4_ZERO
375#define GGML_F32_VEC_SET1   GGML_F32x4_SET1
376#define GGML_F32_VEC_LOAD   GGML_F32x4_LOAD
377#define GGML_F32_VEC_STORE  GGML_F32x4_STORE
378#define GGML_F32_VEC_FMA    GGML_F32x4_FMA
379#define GGML_F32_VEC_ADD    GGML_F32x4_ADD
380#define GGML_F32_VEC_MUL    GGML_F32x4_MUL
381#define GGML_F32_VEC_REDUCE GGML_F32x4_REDUCE
382
383// F16 NEON
384
385#if defined(__ARM_FEATURE_FP16_VECTOR_ARITHMETIC)
386    #define GGML_F16_STEP 32
387    #define GGML_F16_EPR  8
388
389    #define GGML_F16x8              float16x8_t
390    #define GGML_F16x8_ZERO         vdupq_n_f16(0.0f)
391    #define GGML_F16x8_SET1(x)      vdupq_n_f16(x)
392    #define GGML_F16x8_LOAD         vld1q_f16
393    #define GGML_F16x8_STORE        vst1q_f16
394    #define GGML_F16x8_FMA(a, b, c) vfmaq_f16(a, b, c)
395    #define GGML_F16x8_ADD          vaddq_f16
396    #define GGML_F16x8_MUL          vmulq_f16
397    #define GGML_F16x8_REDUCE(res, x)                             \
398    {                                                             \
399        for (int i = 0; i < GGML_F16_ARR/2; ++i) {                \
400            x[2*i] = vaddq_f16(x[2*i], x[2*i+1]);                 \
401        }                                                         \
402        for (int i = 0; i < GGML_F16_ARR/4; ++i) {                \
403            x[4*i] = vaddq_f16(x[4*i], x[4*i+2]);                 \
404        }                                                         \
405        for (int i = 0; i < GGML_F16_ARR/8; ++i) {                \
406            x[8*i] = vaddq_f16(x[8*i], x[8*i+4]);                 \
407        }                                                         \
408        const float32x4_t t0 = vcvt_f32_f16(vget_low_f16 (x[0])); \
409        const float32x4_t t1 = vcvt_f32_f16(vget_high_f16(x[0])); \
410        res = vaddvq_f32(vaddq_f32(t0, t1));                      \
411    }
412
413    #define GGML_F16_VEC        GGML_F16x8
414    #define GGML_F16_VEC_ZERO   GGML_F16x8_ZERO
415    #define GGML_F16_VEC_SET1   GGML_F16x8_SET1
416    #define GGML_F16_VEC_LOAD   GGML_F16x8_LOAD
417    #define GGML_F16_VEC_STORE  GGML_F16x8_STORE
418    #define GGML_F16_VEC_FMA    GGML_F16x8_FMA
419    #define GGML_F16_VEC_ADD    GGML_F16x8_ADD
420    #define GGML_F16_VEC_MUL    GGML_F16x8_MUL
421    #define GGML_F16_VEC_REDUCE GGML_F16x8_REDUCE
422#else
423    // if FP16 vector arithmetic is not supported, we use FP32 instead
424    // and take advantage of the vcvt_ functions to convert to/from FP16
425
426    #define GGML_F16_STEP 16
427    #define GGML_F16_EPR  4
428
429    #define GGML_F32Cx4              float32x4_t
430    #define GGML_F32Cx4_ZERO         vdupq_n_f32(0.0f)
431    #define GGML_F32Cx4_SET1(x)      vdupq_n_f32(x)
432    #define GGML_F32Cx4_LOAD(x)      vcvt_f32_f16(vld1_f16(x))
433    #define GGML_F32Cx4_STORE(x, y)  vst1_f16(x, vcvt_f16_f32(y))
434    #define GGML_F32Cx4_FMA(a, b, c) vfmaq_f32(a, b, c)
435    #define GGML_F32Cx4_ADD          vaddq_f32
436    #define GGML_F32Cx4_MUL          vmulq_f32
437    #define GGML_F32Cx4_REDUCE       GGML_F32x4_REDUCE
438
439    #define GGML_F16_VEC        GGML_F32Cx4
440    #define GGML_F16_VEC_ZERO   GGML_F32Cx4_ZERO
441    #define GGML_F16_VEC_SET1   GGML_F32Cx4_SET1
442    #define GGML_F16_VEC_LOAD   GGML_F32Cx4_LOAD
443    #define GGML_F16_VEC_STORE  GGML_F32Cx4_STORE
444    #define GGML_F16_VEC_FMA    GGML_F32Cx4_FMA
445    #define GGML_F16_VEC_ADD    GGML_F32Cx4_ADD
446    #define GGML_F16_VEC_MUL    GGML_F32Cx4_MUL
447    #define GGML_F16_VEC_REDUCE GGML_F32Cx4_REDUCE
448#endif
449
450#elif defined(__AVX__)
451
452#define GGML_SIMD
453
454// F32 AVX
455
456#define GGML_F32_STEP 32
457#define GGML_F32_EPR  8
458
459#define GGML_F32x8         __m256
460#define GGML_F32x8_ZERO    _mm256_setzero_ps()
461#define GGML_F32x8_SET1(x) _mm256_set1_ps(x)
462#define GGML_F32x8_LOAD    _mm256_loadu_ps
463#define GGML_F32x8_STORE   _mm256_storeu_ps
464#if defined(__FMA__)
465    #define GGML_F32x8_FMA(a, b, c) _mm256_fmadd_ps(b, c, a)
466#else
467    #define GGML_F32x8_FMA(a, b, c) _mm256_add_ps(_mm256_mul_ps(b, c), a)
468#endif
469#define GGML_F32x8_ADD     _mm256_add_ps
470#define GGML_F32x8_MUL     _mm256_mul_ps
471#define GGML_F32x8_REDUCE(res, x)                                 \
472{                                                                 \
473    for (int i = 0; i < GGML_F32_ARR/2; ++i) {                    \
474        x[2*i] = _mm256_add_ps(x[2*i], x[2*i+1]);                 \
475    }                                                             \
476    for (int i = 0; i < GGML_F32_ARR/4; ++i) {                    \
477        x[4*i] = _mm256_add_ps(x[4*i], x[4*i+2]);                 \
478    }                                                             \
479    for (int i = 0; i < GGML_F32_ARR/8; ++i) {                    \
480        x[8*i] = _mm256_add_ps(x[8*i], x[8*i+4]);                 \
481    }                                                             \
482    const __m128 t0 = _mm_add_ps(_mm256_castps256_ps128(x[0]),    \
483                                 _mm256_extractf128_ps(x[0], 1)); \
484    const __m128 t1 = _mm_hadd_ps(t0, t0);                        \
485    res = _mm_cvtss_f32(_mm_hadd_ps(t1, t1));                     \
486}
487// TODO: is this optimal ?
488
489#define GGML_F32_VEC        GGML_F32x8
490#define GGML_F32_VEC_ZERO   GGML_F32x8_ZERO
491#define GGML_F32_VEC_SET1   GGML_F32x8_SET1
492#define GGML_F32_VEC_LOAD   GGML_F32x8_LOAD
493#define GGML_F32_VEC_STORE  GGML_F32x8_STORE
494#define GGML_F32_VEC_FMA    GGML_F32x8_FMA
495#define GGML_F32_VEC_ADD    GGML_F32x8_ADD
496#define GGML_F32_VEC_MUL    GGML_F32x8_MUL
497#define GGML_F32_VEC_REDUCE GGML_F32x8_REDUCE
498
499// F16 AVX
500
501#define GGML_F16_STEP 32
502#define GGML_F16_EPR  8
503
504// F16 arithmetic is not supported by AVX, so we use F32 instead
505// we take advantage of the _mm256_cvt intrinsics to convert F16 <-> F32
506
507#define GGML_F32Cx8             __m256
508#define GGML_F32Cx8_ZERO        _mm256_setzero_ps()
509#define GGML_F32Cx8_SET1(x)     _mm256_set1_ps(x)
510#define GGML_F32Cx8_LOAD(x)     _mm256_cvtph_ps(_mm_loadu_si128((__m128i *)(x)))
511#define GGML_F32Cx8_STORE(x, y) _mm_storeu_si128((__m128i *)(x), _mm256_cvtps_ph(y, 0))
512#define GGML_F32Cx8_FMA         GGML_F32x8_FMA
513#define GGML_F32Cx8_ADD         _mm256_add_ps
514#define GGML_F32Cx8_MUL         _mm256_mul_ps
515#define GGML_F32Cx8_REDUCE      GGML_F32x8_REDUCE
516
517#define GGML_F16_VEC        GGML_F32Cx8
518#define GGML_F16_VEC_ZERO   GGML_F32Cx8_ZERO
519#define GGML_F16_VEC_SET1   GGML_F32Cx8_SET1
520#define GGML_F16_VEC_LOAD   GGML_F32Cx8_LOAD
521#define GGML_F16_VEC_STORE  GGML_F32Cx8_STORE
522#define GGML_F16_VEC_FMA    GGML_F32Cx8_FMA
523#define GGML_F16_VEC_ADD    GGML_F32Cx8_ADD
524#define GGML_F16_VEC_MUL    GGML_F32Cx8_MUL
525#define GGML_F16_VEC_REDUCE GGML_F32Cx8_REDUCE
526
527#elif defined(__POWER9_VECTOR__)
528
529// TODO: uncomment this when it works
530//#define GGML_SIMD
531
532// F32 POWER9
533
534#define GGML_F32_STEP 32
535#define GGML_F32_EPR  8
536
537// TODO: not tested !!
538#define GGML_F32x4         __vector float
539#define GGML_F32x4_ZERO    (__vector float){0.0f, 0.0f, 0.0f, 0.0f}
540#define GGML_F32x4_SET1(x) (__vector float){x, x, x, x}
541#define GGML_F32x4_LOAD    vec_vsx_ld
542#define GGML_F32x4_STORE   vec_vsx_st
543#define GGML_F32x4_FMA(a, b, c) vec_madd(b, c, a)
544#define GGML_F32x4_ADD     vec_add
545#define GGML_F32x4_MUL     vec_mul
546#define GGML_F32x4_REDUCE(res, x)              \
547{                                              \
548    for (int i = 0; i < GGML_F32_ARR/2; ++i) { \
549        x[2*i] = vec_add(x[2*i], x[2*i+1]);    \
550    }                                          \
551    for (int i = 0; i < GGML_F32_ARR/4; ++i) { \
552        x[4*i] = vec_add(x[4*i], x[4*i+2]);    \
553    }                                          \
554    for (int i = 0; i < GGML_F32_ARR/8; ++i) { \
555        x[8*i] = vec_add(x[8*i], x[8*i+4]);    \
556    }                                          \
557    res = vec_extract(x[0], 0) +               \
558          vec_extract(x[0], 1) +               \
559          vec_extract(x[0], 2) +               \
560          vec_extract(x[0], 3);                \
561}
562
563#define GGML_F32_VEC        GGML_F32x4
564#define GGML_F32_VEC_ZERO   GGML_F32x4_ZERO
565#define GGML_F32_VEC_SET1   GGML_F32x4_SET1
566#define GGML_F32_VEC_LOAD   GGML_F32x4_LOAD
567#define GGML_F32_VEC_STORE  GGML_F32x4_STORE
568#define GGML_F32_VEC_FMA    GGML_F32x4_FMA
569#define GGML_F32_VEC_ADD    GGML_F32x4_ADD
570#define GGML_F32_VEC_MUL    GGML_F32x4_MUL
571#define GGML_F32_VEC_REDUCE GGML_F32x4_REDUCE
572
573// F16 POWER9
574// TODO: implement here
575// ...
576
577#elif defined(__wasm_simd128__)
578
579#define GGML_SIMD
580
581// F32 WASM
582
583#define GGML_F32_STEP 16
584#define GGML_F32_EPR  4
585
586#define GGML_F32x4              v128_t
587#define GGML_F32x4_ZERO         wasm_f32x4_splat(0.0f)
588#define GGML_F32x4_SET1(x)      wasm_f32x4_splat(x)
589#define GGML_F32x4_LOAD         wasm_v128_load
590#define GGML_F32x4_STORE        wasm_v128_store
591#define GGML_F32x4_FMA(a, b, c) wasm_f32x4_add(wasm_f32x4_mul(b, c), a)
592#define GGML_F32x4_ADD          wasm_f32x4_add
593#define GGML_F32x4_MUL          wasm_f32x4_mul
594#define GGML_F32x4_REDUCE(res, x)                  \
595{                                                  \
596    for (int i = 0; i < GGML_F32_ARR/2; ++i) {     \
597        x[2*i] = wasm_f32x4_add(x[2*i], x[2*i+1]); \
598    }                                              \
599    for (int i = 0; i < GGML_F32_ARR/4; ++i) {     \
600        x[4*i] = wasm_f32x4_add(x[4*i], x[4*i+2]); \
601    }                                              \
602    for (int i = 0; i < GGML_F32_ARR/8; ++i) {     \
603        x[8*i] = wasm_f32x4_add(x[8*i], x[8*i+4]); \
604    }                                              \
605    res = wasm_f32x4_extract_lane(x[0], 0) +       \
606          wasm_f32x4_extract_lane(x[0], 1) +       \
607          wasm_f32x4_extract_lane(x[0], 2) +       \
608          wasm_f32x4_extract_lane(x[0], 3);        \
609}
610
611#define GGML_F32_VEC        GGML_F32x4
612#define GGML_F32_VEC_ZERO   GGML_F32x4_ZERO
613#define GGML_F32_VEC_SET1   GGML_F32x4_SET1
614#define GGML_F32_VEC_LOAD   GGML_F32x4_LOAD
615#define GGML_F32_VEC_STORE  GGML_F32x4_STORE
616#define GGML_F32_VEC_FMA    GGML_F32x4_FMA
617#define GGML_F32_VEC_ADD    GGML_F32x4_ADD
618#define GGML_F32_VEC_MUL    GGML_F32x4_MUL
619#define GGML_F32_VEC_REDUCE GGML_F32x4_REDUCE
620
621// F16 WASM
622
623#define GGML_F16_STEP 16
624#define GGML_F16_EPR  4
625
626inline static v128_t __wasm_f16x4_load(const ggml_fp16_t * p) {
627    float tmp[4];
628
629    tmp[0] = GGML_FP16_TO_FP32(p[0]);
630    tmp[1] = GGML_FP16_TO_FP32(p[1]);
631    tmp[2] = GGML_FP16_TO_FP32(p[2]);
632    tmp[3] = GGML_FP16_TO_FP32(p[3]);
633
634    return wasm_v128_load(tmp);
635}
636
637inline static void __wasm_f16x4_store(ggml_fp16_t * p, v128_t x) {
638    float tmp[4];
639
640    wasm_v128_store(tmp, x);
641
642    p[0] = GGML_FP32_TO_FP16(tmp[0]);
643    p[1] = GGML_FP32_TO_FP16(tmp[1]);
644    p[2] = GGML_FP32_TO_FP16(tmp[2]);
645    p[3] = GGML_FP32_TO_FP16(tmp[3]);
646}
647
648#define GGML_F16x4             v128_t
649#define GGML_F16x4_ZERO        wasm_f32x4_splat(0.0f)
650#define GGML_F16x4_SET1(x)     wasm_f32x4_splat(x)
651#define GGML_F16x4_LOAD(x)     __wasm_f16x4_load(x)
652#define GGML_F16x4_STORE(x, y) __wasm_f16x4_store(x, y)
653#define GGML_F16x4_FMA         GGML_F32x4_FMA
654#define GGML_F16x4_ADD         wasm_f32x4_add
655#define GGML_F16x4_MUL         wasm_f32x4_mul
656#define GGML_F16x4_REDUCE(res, x)                  \
657{                                                  \
658    for (int i = 0; i < GGML_F16_ARR/2; ++i) {     \
659        x[2*i] = wasm_f32x4_add(x[2*i], x[2*i+1]); \
660    }                                              \
661    for (int i = 0; i < GGML_F16_ARR/4; ++i) {     \
662        x[4*i] = wasm_f32x4_add(x[4*i], x[4*i+2]); \
663    }                                              \
664    for (int i = 0; i < GGML_F16_ARR/8; ++i) {     \
665        x[8*i] = wasm_f32x4_add(x[8*i], x[8*i+4]); \
666    }                                              \
667    res = wasm_f32x4_extract_lane(x[0], 0) +       \
668          wasm_f32x4_extract_lane(x[0], 1) +       \
669          wasm_f32x4_extract_lane(x[0], 2) +       \
670          wasm_f32x4_extract_lane(x[0], 3);        \
671}
672
673#define GGML_F16_VEC        GGML_F16x4
674#define GGML_F16_VEC_ZERO   GGML_F16x4_ZERO
675#define GGML_F16_VEC_SET1   GGML_F16x4_SET1
676#define GGML_F16_VEC_LOAD   GGML_F16x4_LOAD
677#define GGML_F16_VEC_STORE  GGML_F16x4_STORE
678#define GGML_F16_VEC_FMA    GGML_F16x4_FMA
679#define GGML_F16_VEC_ADD    GGML_F16x4_ADD
680#define GGML_F16_VEC_MUL    GGML_F16x4_MUL
681#define GGML_F16_VEC_REDUCE GGML_F16x4_REDUCE
682
683#endif
684
685// GGML_F32_ARR / GGML_F16_ARR
686//   number of registers to use per step
687#ifdef GGML_SIMD
688#define GGML_F32_ARR (GGML_F32_STEP/GGML_F32_EPR)
689#define GGML_F16_ARR (GGML_F16_STEP/GGML_F16_EPR)
690#endif
691
692//
693// fundamental operations
694//
695
696inline static void ggml_vec_set_i8(const int n, int8_t * x, const int8_t v) { for (int i = 0; i < n; ++i) x[i] = v; }
697
698inline static void ggml_vec_set_i16(const int n, int16_t * x, const int16_t v) { for (int i = 0; i < n; ++i) x[i] = v; }
699
700inline static void ggml_vec_set_i32(const int n, int32_t * x, const int32_t v) { for (int i = 0; i < n; ++i) x[i] = v; }
701
702inline static void ggml_vec_set_f16(const int n, ggml_fp16_t * x, const int32_t v) { for (int i = 0; i < n; ++i) x[i] = v; }
703
704inline static void ggml_vec_add_f32 (const int n, float * z, const float * x, const float * y) { for (int i = 0; i < n; ++i) z[i]  = x[i] + y[i]; }
705inline static void ggml_vec_acc_f32 (const int n, float * y, const float * x)                  { for (int i = 0; i < n; ++i) y[i] += x[i];        }
706inline static void ggml_vec_acc1_f32(const int n, float * y, const float   v)                  { for (int i = 0; i < n; ++i) y[i] += v;           }
707inline static void ggml_vec_sub_f32 (const int n, float * z, const float * x, const float * y) { for (int i = 0; i < n; ++i) z[i]  = x[i] - y[i]; }
708inline static void ggml_vec_set_f32 (const int n, float * x, const float   v)                  { for (int i = 0; i < n; ++i) x[i]  = v;           }
709inline static void ggml_vec_cpy_f32 (const int n, float * y, const float * x)                  { for (int i = 0; i < n; ++i) y[i]  = x[i];        }
710inline static void ggml_vec_neg_f32 (const int n, float * y, const float * x)                  { for (int i = 0; i < n; ++i) y[i]  = -x[i];       }
711inline static void ggml_vec_mul_f32 (const int n, float * z, const float * x, const float * y) { for (int i = 0; i < n; ++i) z[i]  = x[i]*y[i];   }
712inline static void ggml_vec_div_f32 (const int n, float * z, const float * x, const float * y) { for (int i = 0; i < n; ++i) z[i]  = x[i]/y[i];   }
713
714inline static void ggml_vec_dot_f32(const int n, float * restrict s, const float * restrict x, const float * restrict y) {
715    ggml_float sumf = 0.0;
716
717#ifdef GGML_SIMD
718    const int np = (n & ~(GGML_F32_STEP - 1));
719
720    GGML_F32_VEC sum[GGML_F32_ARR] = { GGML_F32_VEC_ZERO };
721
722    GGML_F32_VEC ax[GGML_F32_ARR];
723    GGML_F32_VEC ay[GGML_F32_ARR];
724
725    for (int i = 0; i < np; i += GGML_F32_STEP) {
726        for (int j = 0; j < GGML_F32_ARR; j++) {
727            ax[j] = GGML_F32_VEC_LOAD(x + i + j*GGML_F32_EPR);
728            ay[j] = GGML_F32_VEC_LOAD(y + i + j*GGML_F32_EPR);
729
730            sum[j] = GGML_F32_VEC_FMA(sum[j], ax[j], ay[j]);
731        }
732    }
733
734    // reduce sum0..sum3 to sum0
735    GGML_F32_VEC_REDUCE(sumf, sum);
736
737    // leftovers
738    for (int i = np; i < n; ++i) {
739        sumf += x[i]*y[i];
740    }
741#else
742    // scalar
743    for (int i = 0; i < n; ++i) {
744        sumf += x[i]*y[i];
745    }
746#endif
747
748    *s = sumf;
749}
750
751inline static void ggml_vec_dot_f16(const int n, float * restrict s, ggml_fp16_t * restrict x, ggml_fp16_t * restrict y) {
752    ggml_float sumf = 0.0;
753
754#if defined(GGML_SIMD)
755    const int np = (n & ~(GGML_F16_STEP - 1));
756
757    GGML_F16_VEC sum[GGML_F16_ARR] = { GGML_F16_VEC_ZERO };
758
759    GGML_F16_VEC ax[GGML_F16_ARR];
760    GGML_F16_VEC ay[GGML_F16_ARR];
761
762    for (int i = 0; i < np; i += GGML_F16_STEP) {
763        for (int j = 0; j < GGML_F16_ARR; j++) {
764            ax[j] = GGML_F16_VEC_LOAD(x + i + j*GGML_F16_EPR);
765            ay[j] = GGML_F16_VEC_LOAD(y + i + j*GGML_F16_EPR);
766
767            sum[j] = GGML_F16_VEC_FMA(sum[j], ax[j], ay[j]);
768        }
769    }
770
771    // reduce sum0..sum3 to sum0
772    GGML_F16_VEC_REDUCE(sumf, sum);
773
774    // leftovers
775    for (int i = np; i < n; ++i) {
776        sumf += GGML_FP16_TO_FP32(x[i])*GGML_FP16_TO_FP32(y[i]);
777    }
778#elif defined(__POWER9_VECTOR__)
779    // TODO: this is temporary because I cannot fit it in the GGML_SIMD pattern like all other architectures without
780    //       being able to test it. hoping someone with access to a POWER9 machine can help out here.
781    const int n32 = (n & ~31);
782
783    vector float sum0 = vec_splats (0.0f);
784
785    for (int i = 0; i < n32; i += 32) {
786        // Use vec_xl, not vec_ld, because x is sometimes unaligned.
787        vector unsigned short x0 = vec_xl(i * 2 +  0, x);
788        vector unsigned short x1 = vec_xl(i * 2 + 16, x);
789        vector unsigned short x2 = vec_xl(i * 2 + 32, x);
790        vector unsigned short x3 = vec_xl(i * 2 + 48, x);
791
792        vector unsigned short y0 = vec_xl(i * 2 +  0, y);
793        vector unsigned short y1 = vec_xl(i * 2 + 16, y);
794        vector unsigned short y2 = vec_xl(i * 2 + 32, y);
795        vector unsigned short y3 = vec_xl(i * 2 + 48, y);
796
797        vector float fx0l = vec_extract_fp32_from_shortl(x0);
798        vector float fx0h = vec_extract_fp32_from_shorth(x0);
799        vector float fx1l = vec_extract_fp32_from_shortl(x1);
800        vector float fx1h = vec_extract_fp32_from_shorth(x1);
801        vector float fx2l = vec_extract_fp32_from_shortl(x2);
802        vector float fx2h = vec_extract_fp32_from_shorth(x2);
803        vector float fx3l = vec_extract_fp32_from_shortl(x3);
804        vector float fx3h = vec_extract_fp32_from_shorth(x3);
805
806        vector float fy0l = vec_extract_fp32_from_shortl(y0);
807        vector float fy0h = vec_extract_fp32_from_shorth(y0);
808        vector float fy1l = vec_extract_fp32_from_shortl(y1);
809        vector float fy1h = vec_extract_fp32_from_shorth(y1);
810        vector float fy2l = vec_extract_fp32_from_shortl(y2);
811        vector float fy2h = vec_extract_fp32_from_shorth(y2);
812        vector float fy3l = vec_extract_fp32_from_shortl(y3);
813        vector float fy3h = vec_extract_fp32_from_shorth(y3);
814
815        sum0 = vec_add(sum0, vec_mul(fx0l, fy0l));
816        sum0 = vec_add(sum0, vec_mul(fx0h, fy0h));
817        sum0 = vec_add(sum0, vec_mul(fx1l, fy1l));
818        sum0 = vec_add(sum0, vec_mul(fx1h, fy1h));
819        sum0 = vec_add(sum0, vec_mul(fx2l, fy2l));
820        sum0 = vec_add(sum0, vec_mul(fx2h, fy2h));
821        sum0 = vec_add(sum0, vec_mul(fx3l, fy3l));
822        sum0 = vec_add(sum0, vec_mul(fx3h, fy3h));
823    }
824
825    sumf = vec_extract(sum0, 0) + vec_extract(sum0, 1)
826         + vec_extract(sum0, 2) + vec_extract(sum0, 3);
827
828    for (int i = n32; i < n; ++i) {
829        sumf += GGML_FP16_TO_FP32(x[i])*GGML_FP16_TO_FP32(y[i]);
830    }
831#else
832    for (int i = 0; i < n; ++i) {
833        sumf += GGML_FP16_TO_FP32(x[i])*GGML_FP16_TO_FP32(y[i]);
834    }
835#endif
836
837    *s = sumf;
838}
839
840inline static void ggml_vec_mad_f32(const int n, float * restrict y, const float * restrict x, const float v) {
841#if defined(GGML_SIMD)
842    const int np = (n & ~(GGML_F32_STEP - 1));
843
844    GGML_F32_VEC vx = GGML_F32_VEC_SET1(v);
845
846    GGML_F32_VEC ax[GGML_F32_ARR];
847    GGML_F32_VEC ay[GGML_F32_ARR];
848
849    for (int i = 0; i < np; i += GGML_F32_STEP) {
850        for (int j = 0; j < GGML_F32_ARR; j++) {
851            ax[j] = GGML_F32_VEC_LOAD(x + i + j*GGML_F32_EPR);
852            ay[j] = GGML_F32_VEC_LOAD(y + i + j*GGML_F32_EPR);
853            ay[j] = GGML_F32_VEC_FMA(ay[j], ax[j], vx);
854
855            GGML_F32_VEC_STORE(y + i + j*GGML_F32_EPR, ay[j]);
856        }
857    }
858
859    // leftovers
860    for (int i = np; i < n; ++i) {
861        y[i] += x[i]*v;
862    }
863#else
864    // scalar
865    for (int i = 0; i < n; ++i) {
866        y[i] += x[i]*v;
867    }
868#endif
869}
870
871inline static void ggml_vec_mad_f16(const int n, ggml_fp16_t * restrict y, ggml_fp16_t * restrict x, const float v) {
872#if defined(GGML_SIMD)
873    const int np = (n & ~(GGML_F16_STEP - 1));
874
875    GGML_F16_VEC vx = GGML_F16_VEC_SET1(v);
876
877    GGML_F16_VEC ax[GGML_F16_ARR];
878    GGML_F16_VEC ay[GGML_F16_ARR];
879
880    for (int i = 0; i < np; i += GGML_F16_STEP) {
881        for (int j = 0; j < GGML_F16_ARR; j++) {
882            ax[j] = GGML_F16_VEC_LOAD(x + i + j*GGML_F16_EPR);
883            ay[j] = GGML_F16_VEC_LOAD(y + i + j*GGML_F16_EPR);
884            ay[j] = GGML_F16_VEC_FMA(ay[j], ax[j], vx);
885
886            GGML_F16_VEC_STORE(y + i + j*GGML_F16_EPR, ay[j]);
887        }
888    }
889
890    // leftovers
891    for (int i = np; i < n; ++i) {
892        GGML_ASSERT(false);
893        y[i] = GGML_FP32_TO_FP16(GGML_FP16_TO_FP32(y[i]) + GGML_FP16_TO_FP32(x[i])*v);
894    }
895#elif defined(__POWER9_VECTOR__)
896    // TODO: this is temporary because I cannot fit it in the GGML_SIMD pattern like all other architectures without
897    //       being able to test it. hoping someone with access to a POWER9 machine can help out here.
898    const int n32 = (n & ~31);
899    for (int i = 0; i < n32; i += 32) {
900        // Use vec_xl, not vec_ld, because x is sometimes unaligned!
901        vector unsigned short x0 = vec_xl(i * 2 +  0, x);
902        vector unsigned short x1 = vec_xl(i * 2 + 16, x);
903        vector unsigned short x2 = vec_xl(i * 2 + 32, x);
904        vector unsigned short x3 = vec_xl(i * 2 + 48, x);
905
906        vector unsigned short y0 = vec_xl(i * 2 +  0, y);
907        vector unsigned short y1 = vec_xl(i * 2 + 16, y);
908        vector unsigned short y2 = vec_xl(i * 2 + 32, y);
909        vector unsigned short y3 = vec_xl(i * 2 + 48, y);
910
911        vector float v4 = vec_splats(v);
912
913        vector float fx0l = vec_extract_fp32_from_shortl(x0);
914        vector float fx0h = vec_extract_fp32_from_shorth(x0);
915        vector float fx1l = vec_extract_fp32_from_shortl(x1);
916        vector float fx1h = vec_extract_fp32_from_shorth(x1);
917        vector float fx2l = vec_extract_fp32_from_shortl(x2);
918        vector float fx2h = vec_extract_fp32_from_shorth(x2);
919        vector float fx3l = vec_extract_fp32_from_shortl(x3);
920        vector float fx3h = vec_extract_fp32_from_shorth(x3);
921
922        vector float fy0l = vec_extract_fp32_from_shortl(y0);
923        vector float fy0h = vec_extract_fp32_from_shorth(y0);
924        vector float fy1l = vec_extract_fp32_from_shortl(y1);
925        vector float fy1h = vec_extract_fp32_from_shorth(y1);
926        vector float fy2l = vec_extract_fp32_from_shortl(y2);
927        vector float fy2h = vec_extract_fp32_from_shorth(y2);
928        vector float fy3l = vec_extract_fp32_from_shortl(y3);
929        vector float fy3h = vec_extract_fp32_from_shorth(y3);
930
931        fy0l = vec_madd(fx0l, v4, fy0l);
932        fy0h = vec_madd(fx0h, v4, fy0h);
933        fy1l = vec_madd(fx1l, v4, fy1l);
934        fy1h = vec_madd(fx1h, v4, fy1h);
935        fy2l = vec_madd(fx2l, v4, fy2l);
936        fy2h = vec_madd(fx2h, v4, fy2h);
937        fy3l = vec_madd(fx3l, v4, fy3l);
938        fy3h = vec_madd(fx3h, v4, fy3h);
939
940        y0 = vec_pack_to_short_fp32(fy0h, fy0l);
941        y1 = vec_pack_to_short_fp32(fy1h, fy1l);
942        y2 = vec_pack_to_short_fp32(fy2h, fy2l);
943        y3 = vec_pack_to_short_fp32(fy3h, fy3l);
944
945        vec_xst(y0, i * 2 +  0, y);
946        vec_xst(y1, i * 2 + 16, y);
947        vec_xst(y2, i * 2 + 32, y);
948        vec_xst(y3, i * 2 + 48, y);
949    }
950
951    for (int i = n32; i < n; ++i) {
952        y[i] = GGML_FP32_TO_FP16(GGML_FP16_TO_FP32(y[i]) + GGML_FP16_TO_FP32(x[i])*v);
953    }
954#else
955    for (int i = 0; i < n; ++i) {
956        y[i] = GGML_FP32_TO_FP16(GGML_FP16_TO_FP32(y[i]) + GGML_FP16_TO_FP32(x[i])*v);
957    }
958#endif
959}
960
961//inline static void ggml_vec_scale_f32(const int n, float * y, const float   v) { for (int i = 0; i < n; ++i) y[i] *= v;          }
962inline static void ggml_vec_scale_f32(const int n, float * y, const float   v) {
963#if defined(GGML_SIMD)
964    const int np = (n & ~(GGML_F32_STEP - 1));
965
966    GGML_F32_VEC vx = GGML_F32_VEC_SET1(v);
967
968    GGML_F32_VEC ay[GGML_F32_ARR];
969
970    for (int i = 0; i < np; i += GGML_F32_STEP) {
971        for (int j = 0; j < GGML_F32_ARR; j++) {
972            ay[j] = GGML_F32_VEC_LOAD(y + i + j*GGML_F32_EPR);
973            ay[j] = GGML_F32_VEC_MUL(ay[j], vx);
974
975            GGML_F32_VEC_STORE(y + i + j*GGML_F32_EPR, ay[j]);
976        }
977    }
978
979    // leftovers
980    for (int i = np; i < n; ++i) {
981        y[i] *= v;
982    }
983#else
984    // scalar
985    for (int i = 0; i < n; ++i) {
986        y[i] *= v;
987    }
988#endif
989}
990
991inline static void ggml_vec_norm_f32 (const int n, float * s, const float * x) { ggml_vec_dot_f32(n, s, x, x); *s = sqrt(*s);   }
992inline static void ggml_vec_sqr_f32  (const int n, float * y, const float * x) { for (int i = 0; i < n; ++i) y[i] = x[i]*x[i];   }
993inline static void ggml_vec_sqrt_f32 (const int n, float * y, const float * x) { for (int i = 0; i < n; ++i) y[i] = sqrt(x[i]); }
994inline static void ggml_vec_abs_f32  (const int n, float * y, const float * x) { for (int i = 0; i < n; ++i) y[i] = fabsf(x[i]); }
995inline static void ggml_vec_sgn_f32  (const int n, float * y, const float * x) { for (int i = 0; i < n; ++i) y[i] = (x[i] > 0.f) ? 1.f : ((x[i] < 0.f) ? -1.f : 0.f); }
996inline static void ggml_vec_step_f32 (const int n, float * y, const float * x) { for (int i = 0; i < n; ++i) y[i] = (x[i] > 0.f) ? 1.f : 0.f; }
997inline static void ggml_vec_relu_f32 (const int n, float * y, const float * x) { for (int i = 0; i < n; ++i) y[i] = (x[i] > 0.f) ? x[i] : 0.f; }
998
999static const ggml_float GELU_COEF_A    = 0.044715;
1000static const ggml_float SQRT_2_OVER_PI = 0.79788456080286535587989211986876;
1001
1002inline static float ggml_gelu_f32(float x) {
1003    return 0.5*x*(1.0 + tanh(SQRT_2_OVER_PI*x*(1.0 + GELU_COEF_A*x*x)));
1004}
1005
1006inline static void ggml_vec_gelu_f16(const int n, ggml_fp16_t * y, const ggml_fp16_t * x) {
1007    const uint16_t * i16 = (const uint16_t *) x;
1008    for (int i = 0; i < n; ++i) {
1009        y[i] = table_gelu_f16[i16[i]];
1010    }
1011}
1012
1013#ifdef GGML_GELU_FP16
1014inline static void ggml_vec_gelu_f32(const int n, float * y, const float * x) {
1015    uint16_t t;
1016    for (int i = 0; i < n; ++i) {
1017        ggml_fp16_t fp16 = GGML_FP32_TO_FP16(x[i]);
1018        memcpy(&t, &fp16, sizeof(uint16_t));
1019        y[i] = GGML_FP16_TO_FP32(table_gelu_f16[t]);
1020    }
1021}
1022#else
1023inline static void ggml_vec_gelu_f32(const int n, float * y, const float * x) {
1024    for (int i = 0; i < n; ++i) {
1025        y[i] = ggml_gelu_f32(x[i]);
1026    }
1027}
1028#endif
1029
1030inline static void ggml_vec_sum_f32     (const int n, float * s, const float * x) { ggml_float sum = 0.0; for (int i = 0; i < n; ++i) sum += x[i]; *s += sum; }
1031inline static void ggml_vec_norm_inv_f32(const int n, float * s, const float * x) { ggml_vec_norm_f32(n, s, x); *s = 1./(*s); }
1032
1033//
1034// logging
1035//
1036
1037#if (GGML_DEBUG >= 1)
1038#define GGML_PRINT_DEBUG(...) printf(__VA_ARGS__)
1039#else
1040#define GGML_PRINT_DEBUG(...)
1041#endif
1042
1043#if (GGML_DEBUG >= 5)
1044#define GGML_PRINT_DEBUG_5(...) printf(__VA_ARGS__)
1045#else
1046#define GGML_PRINT_DEBUG_5(...)
1047#endif
1048
1049#if (GGML_DEBUG >= 10)
1050#define GGML_PRINT_DEBUG_10(...) printf(__VA_ARGS__)
1051#else
1052#define GGML_PRINT_DEBUG_10(...)
1053#endif
1054
1055#define GGML_PRINT(...) logDebug( __VA_ARGS__ )
1056
1057//
1058// data types
1059//
1060
1061static const size_t GGML_TYPE_SIZE[GGML_TYPE_COUNT] = {
1062    sizeof(int8_t ),
1063    sizeof(int16_t),
1064    sizeof(int32_t),
1065    sizeof(ggml_fp16_t),
1066    sizeof(float  ),
1067};
1068
1069static const char * GGML_OP_LABEL[GGML_OP_COUNT] = {
1070    "NONE",
1071
1072    "DUP",
1073    "ADD",
1074    "SUB",
1075    "MUL",
1076    "DIV",
1077    "SQR",
1078    "SQRT",
1079    "SUM",
1080    "MEAN",
1081    "REPEAT",
1082    "ABS",
1083    "SGN",
1084    "NEG",
1085    "STEP",
1086    "RELU",
1087    "GELU",
1088    "NORM",
1089
1090    "MUL_MAT",
1091
1092    "SCALE",
1093    "CPY",
1094    "RESHAPE",
1095    "VIEW",
1096    "PERMUTE",
1097    "TRANSPOSE",
1098    "GET_ROWS",
1099    "DIAG_MASK_INF",
1100    "SOFT_MAX",
1101    "ROPE",
1102    "CONV_1D_1S",
1103    "CONV_1D_2S",
1104
1105    "FLASH_ATTN",
1106    "FLASH_FF",
1107};
1108
1109static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = {
1110    "none",
1111
1112    "x",
1113    "x+y",
1114    "x-y",
1115    "x*y",
1116    "x/y",
1117    "x^2",
1118    "√x",
1119    "Σx",
1120    "Σx/n",
1121    "repeat(x)",
1122    "abs(x)",
1123    "sgn(x)",
1124    "-x",
1125    "step(x)",
1126    "relu(x)",
1127    "gelu(x)",
1128    "norm(x)",
1129
1130    "X*Y",
1131
1132    "x*v",
1133    "x-\\>y",
1134    "reshape(x)",
1135    "view(x)",
1136    "permute(x)",
1137    "transpose(x)",
1138    "get_rows(x)",
1139    "diag_mask_inf(x)",
1140    "soft_max(x)",
1141    "rope(x)",
1142    "conv_1d_1s(x)",
1143    "conv_1d_2s(x)",
1144
1145    "flash_attn(x)",
1146    "flash_ff(x)",
1147};
1148
1149//
1150// ggml object
1151//
1152
1153struct ggml_object {
1154    size_t offset;
1155    size_t size;
1156
1157    struct ggml_object * next;
1158
1159    char padding[8];
1160};
1161
1162static const size_t GGML_OBJECT_SIZE = sizeof(struct ggml_object);
1163
1164static_assert(sizeof(struct ggml_object)%GGML_MEM_ALIGN == 0, "ggml_object size must be a multiple of GGML_MEM_ALIGN");
1165static_assert(sizeof(struct ggml_tensor)%GGML_MEM_ALIGN == 0, "ggml_tensor size must be a multiple of GGML_MEM_ALIGN");
1166
1167//
1168// ggml context
1169//
1170
1171struct ggml_context {
1172    size_t mem_size;
1173    void * mem_buffer;
1174    bool   mem_buffer_owned;
1175
1176    int n_objects;
1177
1178    struct ggml_object * objects_begin;
1179    struct ggml_object * objects_end;
1180};
1181
1182struct ggml_context_container {
1183    bool used;
1184
1185    struct ggml_context context;
1186};
1187
1188//
1189// compute types
1190//
1191
1192enum ggml_task_type {
1193    GGML_TASK_INIT = 0,
1194    GGML_TASK_COMPUTE,
1195    GGML_TASK_FINALIZE,
1196};
1197
1198struct ggml_compute_params {
1199    enum ggml_task_type type;
1200
1201    int ith, nth;
1202
1203    // work buffer for all threads
1204    size_t wsize;
1205    void * wdata;
1206};
1207
1208//
1209// ggml state
1210//
1211
1212struct ggml_state {
1213    struct ggml_context_container contexts[GGML_MAX_CONTEXTS];
1214};
1215
1216// global state
1217static struct ggml_state g_state;
1218static atomic_int g_state_barrier = 0;
1219
1220// barrier via spin lock
1221inline static void ggml_critical_section_start() {
1222    int processing = atomic_fetch_add(&g_state_barrier, 1);
1223
1224    while (processing > 0) {
1225        // wait for other threads to finish
1226        atomic_fetch_sub(&g_state_barrier, 1);
1227        sched_yield(); // TODO: reconsider this
1228        processing = atomic_fetch_add(&g_state_barrier, 1);
1229    }
1230}
1231
1232// TODO: make this somehow automatically executed
1233//       some sort of "sentry" mechanism
1234inline static void ggml_critical_section_end() {
1235    atomic_fetch_sub(&g_state_barrier, 1);
1236}
1237
1238////////////////////////////////////////////////////////////////////////////////
1239
1240void ggml_print_object(const struct ggml_object * obj) {
1241    GGML_PRINT(" - ggml_object: offset = %zu, size = %zu, next = %p\n",
1242            obj->offset, obj->size, (const void *) obj->next);
1243}
1244
1245void ggml_print_objects(const struct ggml_context * ctx) {
1246    struct ggml_object * obj = ctx->objects_begin;
1247
1248    GGML_PRINT("%s: objects in context %p:\n", __func__, (const void *) ctx);
1249
1250    while (obj != NULL) {
1251        ggml_print_object(obj);
1252        obj = obj->next;
1253    }
1254
1255    GGML_PRINT("%s: --- end ---\n", __func__);
1256}
1257
1258int ggml_nelements(const struct ggml_tensor * tensor) {
1259    static_assert(GGML_MAX_DIMS == 4, "GGML_MAX_DIMS is not 4 - update this function");
1260
1261    return tensor->ne[0]*tensor->ne[1]*tensor->ne[2]*tensor->ne[3];
1262}
1263
1264int ggml_nrows(const struct ggml_tensor * tensor) {
1265    static_assert(GGML_MAX_DIMS == 4, "GGML_MAX_DIMS is not 4 - update this function");
1266
1267    return tensor->ne[1]*tensor->ne[2]*tensor->ne[3];
1268}
1269
1270size_t ggml_nbytes(const struct ggml_tensor * tensor) {
1271    static_assert(GGML_MAX_DIMS == 4, "GGML_MAX_DIMS is not 4 - update this function");
1272
1273    return ggml_nelements(tensor)*GGML_TYPE_SIZE[tensor->type];
1274}
1275
1276size_t ggml_type_size(enum ggml_type type) {
1277    return GGML_TYPE_SIZE[type];
1278}
1279
1280size_t ggml_element_size(const struct ggml_tensor * tensor) {
1281    return GGML_TYPE_SIZE[tensor->type];
1282}
1283
1284bool ggml_is_scalar(const struct ggml_tensor * tensor) {
1285    static_assert(GGML_MAX_DIMS == 4, "GGML_MAX_DIMS is not 4 - update this function");
1286
1287    return tensor->ne[0] == 1 && tensor->ne[1] == 1 && tensor->ne[2] == 1 && tensor->ne[3] == 1;
1288}
1289
1290bool ggml_is_vector(const struct ggml_tensor * tensor) {
1291    static_assert(GGML_MAX_DIMS == 4, "GGML_MAX_DIMS is not 4 - update this function");
1292
1293    return tensor->ne[1] == 1 && tensor->ne[2] == 1 && tensor->ne[3] == 1;
1294}
1295
1296bool ggml_is_matrix(const struct ggml_tensor * tensor) {
1297    static_assert(GGML_MAX_DIMS == 4, "GGML_MAX_DIMS is not 4 - update this function");
1298
1299    return tensor->ne[2] == 1 && tensor->ne[3] == 1;
1300}
1301
1302bool ggml_can_mul_mat(const struct ggml_tensor * t0, const struct ggml_tensor * t1) {
1303    static_assert(GGML_MAX_DIMS == 4, "GGML_MAX_DIMS is not 4 - update this function");
1304
1305    return
1306        (t0->ne[0]  == t1->ne[0])  &&
1307        (t0->ne[2]  == t1->ne[2])  &&
1308        (t0->ne[3]  == t1->ne[3]);
1309}
1310
1311bool ggml_is_contiguous(const struct ggml_tensor * tensor) {
1312    static_assert(GGML_MAX_DIMS == 4, "GGML_MAX_DIMS is not 4 - update this function");
1313
1314    return
1315        tensor->nb[0] == GGML_TYPE_SIZE[tensor->type] &&
1316        tensor->nb[1] == tensor->nb[0]*tensor->ne[0] &&
1317        tensor->nb[2] == tensor->nb[1]*tensor->ne[1] &&
1318        tensor->nb[3] == tensor->nb[2]*tensor->ne[2];
1319}
1320
1321bool ggml_is_padded_1d(const struct ggml_tensor * tensor) {
1322    static_assert(GGML_MAX_DIMS == 4, "GGML_MAX_DIMS is not 4 - update this function");
1323
1324    return
1325        tensor->nb[0] == GGML_TYPE_SIZE[tensor->type] &&
1326        tensor->nb[2] == tensor->nb[1]*tensor->ne[1] &&
1327        tensor->nb[3] == tensor->nb[2]*tensor->ne[2];
1328}
1329
1330bool ggml_are_same_shape(const struct ggml_tensor * t0, const struct ggml_tensor * t1) {
1331    static_assert(GGML_MAX_DIMS == 4, "GGML_MAX_DIMS is not 4 - update this function");
1332
1333    return
1334        (t0->ne[0] == t1->ne[0] ) &&
1335        (t0->ne[1] == t1->ne[1] ) &&
1336        (t0->ne[2] == t1->ne[2] ) &&
1337        (t0->ne[3] == t1->ne[3] );
1338}
1339
1340// check if t1 can be represented as a repeatition of t0
1341bool ggml_can_repeat(const struct ggml_tensor * t0, const struct ggml_tensor * t1) {
1342    static_assert(GGML_MAX_DIMS == 4, "GGML_MAX_DIMS is not 4 - update this function");
1343
1344    return
1345        (t1->ne[0]%t0->ne[0] == 0) &&
1346        (t1->ne[1]%t0->ne[1] == 0) &&
1347        (t1->ne[2]%t0->ne[2] == 0) &&
1348        (t1->ne[3]%t0->ne[3] == 0);
1349}
1350
1351int ggml_up32(int n) {
1352    return (n + 31) & ~31;
1353}
1354
1355int ggml_up64(int n) {
1356    return (n + 63) & ~63;
1357}
1358
1359// assert that pointer is aligned to GGML_MEM_ALIGN
1360#define ggml_assert_aligned(ptr) \
1361    assert(((uintptr_t) (ptr))%GGML_MEM_ALIGN == 0)
1362
1363////////////////////////////////////////////////////////////////////////////////
1364
1365struct ggml_context * ggml_init(struct ggml_init_params params) {
1366    // make this function thread safe
1367    ggml_critical_section_start();
1368
1369    static bool is_first_call = true;
1370
1371    if (is_first_call) {
1372        // initialize GELU and EXP tables
1373        {
1374            const uint64_t t_start = ggml_time_us(); UNUSED(t_start);
1375
1376            ggml_fp16_t ii;
1377            for (int i = 0; i < (1 << 16); ++i) {
1378                uint16_t ui = i;
1379                memcpy(&ii, &ui, sizeof(ii));
1380                const float f = GGML_FP16_TO_FP32(ii);
1381                table_gelu_f16[i] = GGML_FP32_TO_FP16(ggml_gelu_f32(f));
1382                table_exp_f16[i]  = GGML_FP32_TO_FP16(exp(f));
1383            }
1384
1385            const uint64_t t_end = ggml_time_us(); UNUSED(t_end);
1386
1387            GGML_PRINT_DEBUG("%s: GELU and EXP tables initialized in %f ms\n", __func__, (t_end - t_start)/1000.0f);
1388        }
1389
1390        // initialize g_state
1391        {
1392            const uint64_t t_start = ggml_time_us(); UNUSED(t_start);
1393
1394            g_state = (struct ggml_state) {
1395                /*.contexts =*/ { 0 },
1396            };
1397
1398            for (int i = 0; i < GGML_MAX_CONTEXTS; ++i) {
1399                g_state.contexts[i].used = false;
1400            }
1401
1402            const uint64_t t_end = ggml_time_us(); UNUSED(t_end);
1403
1404            GGML_PRINT_DEBUG("%s: g_state initialized in %f ms\n", __func__, (t_end - t_start)/1000.0f);
1405        }
1406
1407        is_first_call = false;
1408    }
1409
1410    // find non-used context in g_state
1411    struct ggml_context * ctx = NULL;
1412
1413    for (int i = 0; i < GGML_MAX_CONTEXTS; i++) {
1414        if (!g_state.contexts[i].used) {
1415            g_state.contexts[i].used = true;
1416            ctx = &g_state.contexts[i].context;
1417
1418            GGML_PRINT_DEBUG("%s: found unused context %d\n", __func__, i);
1419            break;
1420        }
1421    }
1422
1423    if (ctx == NULL) {
1424        GGML_PRINT_DEBUG("%s: no unused context found\n", __func__);
1425
1426        ggml_critical_section_end();
1427
1428        return NULL;
1429    }
1430
1431    *ctx = (struct ggml_context) {
1432        .mem_size         = params.mem_size,
1433        .mem_buffer       = params.mem_buffer ? params.mem_buffer : malloc(params.mem_size),
1434        .mem_buffer_owned = params.mem_buffer ? false : true,
1435        .n_objects        = 0,
1436        .objects_begin    = NULL,
1437        .objects_end      = NULL,
1438    };
1439
1440    ggml_assert_aligned(ctx->mem_buffer);
1441
1442    GGML_PRINT_DEBUG("%s: context initialized\n", __func__);
1443
1444    ggml_critical_section_end();
1445
1446    return ctx;
1447}
1448
1449void ggml_free(struct ggml_context * ctx) {
1450    // make this function thread safe
1451    ggml_critical_section_start();
1452
1453    bool found = false;
1454
1455    for (int i = 0; i < GGML_MAX_CONTEXTS; i++) {
1456        if (&g_state.contexts[i].context == ctx) {
1457            g_state.contexts[i].used = false;
1458
1459            GGML_PRINT_DEBUG("%s: context %d with %d objects has been freed. memory used = %zu\n",
1460                    __func__, i, ctx->n_objects, ctx->objects_end->offset + ctx->objects_end->size);
1461
1462            if (ctx->mem_buffer_owned) {
1463                free(ctx->mem_buffer);
1464            }
1465
1466            found = true;
1467            break;
1468        }
1469    }
1470
1471    if (!found) {
1472        GGML_PRINT_DEBUG("%s: context not found\n", __func__);
1473    }
1474
1475    ggml_critical_section_end();
1476}
1477
1478size_t ggml_used_mem(const struct ggml_context * ctx) {
1479    return ctx->objects_end->offset + ctx->objects_end->size;
1480}
1481
1482////////////////////////////////////////////////////////////////////////////////
1483
1484struct ggml_tensor * ggml_new_tensor_impl(
1485        struct ggml_context * ctx,
1486        enum   ggml_type type,
1487        int    n_dims,
1488        const int* ne,
1489        void*  data) {
1490    // always insert objects at the end of the context's memory pool
1491    struct ggml_object * obj_cur = ctx->objects_end;
1492
1493    const size_t cur_offset = obj_cur == NULL ? 0 : obj_cur->offset;
1494    const size_t cur_size   = obj_cur == NULL ? 0 : obj_cur->size;
1495    const size_t cur_end    = cur_offset + cur_size;
1496
1497    size_t size_needed = 0;
1498
1499    if (data == NULL) {
1500        size_needed += GGML_TYPE_SIZE[type];
1501        for (int i = 0; i < n_dims; i++) {
1502            size_needed *= ne[i];
1503        }
1504        // align to GGML_MEM_ALIGN
1505        size_needed = ((size_needed + GGML_MEM_ALIGN - 1)/GGML_MEM_ALIGN)*GGML_MEM_ALIGN;
1506
1507    }
1508    size_needed += sizeof(struct ggml_tensor);
1509
1510    if (cur_end + size_needed + GGML_OBJECT_SIZE > ctx->mem_size) {
1511        GGML_PRINT("%s: not enough space in the context's memory pool\n", __func__);
1512        assert(false);
1513        return NULL;
1514    }
1515
1516    char * const mem_buffer = ctx->mem_buffer;
1517
1518    struct ggml_object * const obj_new = (struct ggml_object *)(mem_buffer + cur_end);
1519
1520    *obj_new = (struct ggml_object) {
1521        .offset = cur_end + GGML_OBJECT_SIZE,
1522        .size   = size_needed,
1523        .next   = NULL,
1524    };
1525
1526    if (obj_cur != NULL) {
1527        obj_cur->next = obj_new;
1528    } else {
1529        // this is the first object in this context
1530        ctx->objects_begin = obj_new;
1531    }
1532
1533    ctx->objects_end = obj_new;
1534
1535    //GGML_PRINT_DEBUG("%s: inserted new object at %zu\n", __func__, cur_end);
1536
1537    struct ggml_tensor * const result = (struct ggml_tensor *)(mem_buffer + obj_new->offset);
1538
1539    ggml_assert_aligned(result);
1540
1541    *result = (struct ggml_tensor) {
1542        /*.type         =*/ type,
1543        /*.n_dims       =*/ n_dims,
1544        /*.ne           =*/ { 1, 1, 1, 1 },
1545        /*.nb           =*/ { 0, 0, 0, 0 },
1546        /*.op           =*/ GGML_OP_NONE,
1547        /*.is_param     =*/ false,
1548        /*.grad         =*/ NULL,
1549        /*.src0         =*/ NULL,
1550        /*.src1         =*/ NULL,
1551        /*.opt          =*/ { NULL },
1552        /*.n_tasks      =*/ 0,
1553        /*.perf_runs    =*/ 0,
1554        /*.perf_cycles  =*/ 0,
1555        /*.perf_time_us =*/ 0,
1556        /*.data         =*/ data == NULL ? (void *)(result + 1) : data,
1557        /*.pad          =*/ { 0 },
1558    };
1559
1560    ggml_assert_aligned(result->data);
1561
1562    for (int i = 0; i < n_dims; i++) {
1563        result->ne[i] = ne[i];
1564    }
1565
1566    result->nb[0] = GGML_TYPE_SIZE[type];
1567    for (int i = 1; i < GGML_MAX_DIMS; i++) {
1568        result->nb[i] = result->nb[i - 1]*result->ne[i - 1];
1569    }
1570
1571    ctx->n_objects++;
1572
1573    return result;
1574}
1575
1576struct ggml_tensor * ggml_new_tensor(
1577        struct ggml_context * ctx,
1578        enum   ggml_type type,
1579        int    n_dims,
1580        const int* ne) {
1581    return ggml_new_tensor_impl(ctx, type, n_dims, ne, NULL);
1582}
1583
1584struct ggml_tensor * ggml_new_tensor_1d(
1585        struct ggml_context * ctx,
1586        enum   ggml_type type,
1587        int    ne0) {
1588    return ggml_new_tensor(ctx, type, 1, &ne0);
1589}
1590
1591struct ggml_tensor * ggml_new_tensor_2d(
1592        struct ggml_context * ctx,
1593        enum   ggml_type type,
1594        int    ne0,
1595        int    ne1) {
1596    const int ne[2] = { ne0, ne1 };
1597    return ggml_new_tensor(ctx, type, 2, ne);
1598}
1599
1600struct ggml_tensor * ggml_new_tensor_3d(
1601        struct ggml_context * ctx,
1602        enum   ggml_type type,
1603        int    ne0,
1604        int    ne1,
1605        int    ne2) {
1606    const int ne[3] = { ne0, ne1, ne2 };
1607    return ggml_new_tensor(ctx, type, 3, ne);
1608}
1609
1610struct ggml_tensor * ggml_new_tensor_4d(
1611        struct ggml_context * ctx,
1612        enum   ggml_type type,
1613        int    ne0,
1614        int    ne1,
1615        int    ne2,
1616        int    ne3) {
1617    const int ne[4] = { ne0, ne1, ne2, ne3 };
1618    return ggml_new_tensor(ctx, type, 4, ne);
1619}
1620
1621struct ggml_tensor * ggml_new_i32(struct ggml_context * ctx, int32_t value) {
1622    struct ggml_tensor * result = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1);
1623
1624    ggml_set_i32(result, value);
1625
1626    return result;
1627}
1628
1629struct ggml_tensor * ggml_new_f32(struct ggml_context * ctx, float value) {
1630    struct ggml_tensor * result = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 1);
1631
1632    ggml_set_f32(result, value);
1633
1634    return result;
1635}
1636
1637struct ggml_tensor * ggml_dup_tensor(struct ggml_context * ctx, const struct ggml_tensor * src) {
1638    return ggml_new_tensor_impl(ctx, src->type, src->n_dims, src->ne, NULL);
1639}
1640
1641struct ggml_tensor * ggml_set_zero(struct ggml_tensor * tensor) {
1642    memset(tensor->data, 0, ggml_nbytes(tensor));
1643    return tensor;
1644}
1645
1646struct ggml_tensor * ggml_set_i32 (struct ggml_tensor * tensor, int32_t value) {
1647    const int n     = ggml_nrows(tensor);
1648    const int nc    = tensor->ne[0];
1649    const size_t n1 = tensor->nb[1];
1650
1651    char * const data = tensor->data;
1652
1653    switch (tensor->type) {
1654        case GGML_TYPE_I8:
1655            {
1656                assert(tensor->nb[0] == sizeof(int8_t));
1657                for (int i = 0; i < n; i++) {
1658                    ggml_vec_set_i8(nc, (int8_t *)(data + i*n1), value);
1659                }
1660            } break;
1661        case GGML_TYPE_I16:
1662            {
1663                assert(tensor->nb[0] == sizeof(int16_t));
1664                for (int i = 0; i < n; i++) {
1665                    ggml_vec_set_i16(nc, (int16_t *)(data + i*n1), value);
1666                }
1667            } break;
1668        case GGML_TYPE_I32:
1669            {
1670                assert(tensor->nb[0] == sizeof(int32_t));
1671                for (int i = 0; i < n; i++) {
1672                    ggml_vec_set_i32(nc, (int32_t *)(data + i*n1), value);
1673                }
1674            } break;
1675        case GGML_TYPE_F16:
1676            {
1677                assert(tensor->nb[0] == sizeof(ggml_fp16_t));
1678                for (int i = 0; i < n; i++) {
1679                    ggml_vec_set_f16(nc, (ggml_fp16_t *)(data + i*n1), value);
1680                }
1681            } break;
1682        case GGML_TYPE_F32:
1683            {
1684                assert(tensor->nb[0] == sizeof(float));
1685                for (int i = 0; i < n; i++) {
1686                    ggml_vec_set_f32(nc, (float *)(data + i*n1), value);
1687                }
1688            } break;
1689        case GGML_TYPE_COUNT:
1690            {
1691                assert(false);
1692            } break;
1693    }
1694
1695    return tensor;
1696}
1697
1698struct ggml_tensor * ggml_set_f32(struct ggml_tensor * tensor, float value) {
1699    const int n     = ggml_nrows(tensor);
1700    const int nc    = tensor->ne[0];
1701    const size_t n1 = tensor->nb[1];
1702
1703    char * const data = tensor->data;
1704
1705    switch (tensor->type) {
1706        case GGML_TYPE_I8:
1707            {
1708                assert(tensor->nb[0] == sizeof(int8_t));
1709                for (int i = 0; i < n; i++) {
1710                    ggml_vec_set_i8(nc, (int8_t *)(data + i*n1), value);
1711                }
1712            } break;
1713        case GGML_TYPE_I16:
1714            {
1715                assert(tensor->nb[0] == sizeof(int16_t));
1716                for (int i = 0; i < n; i++) {
1717                    ggml_vec_set_i16(nc, (int16_t *)(data + i*n1), value);
1718                }
1719            } break;
1720        case GGML_TYPE_I32:
1721            {
1722                assert(tensor->nb[0] == sizeof(int32_t));
1723                for (int i = 0; i < n; i++) {
1724                    ggml_vec_set_i32(nc, (int32_t *)(data + i*n1), value);
1725                }
1726            } break;
1727        case GGML_TYPE_F16:
1728            {
1729                assert(tensor->nb[0] == sizeof(ggml_fp16_t));
1730                for (int i = 0; i < n; i++) {
1731                    ggml_vec_set_f16(nc, (ggml_fp16_t *)(data + i*n1), value);
1732                }
1733            } break;
1734        case GGML_TYPE_F32:
1735            {
1736                assert(tensor->nb[0] == sizeof(float));
1737                for (int i = 0; i < n; i++) {
1738                    ggml_vec_set_f32(nc, (float *)(data + i*n1), value);
1739                }
1740            } break;
1741        case GGML_TYPE_COUNT:
1742            {
1743                assert(false);
1744            } break;
1745    }
1746
1747    return tensor;
1748}
1749
1750int32_t ggml_get_i32_1d(const struct ggml_tensor * tensor, int i) {
1751    switch (tensor->type) {
1752        case GGML_TYPE_I8:
1753            {
1754                GGML_ASSERT(tensor->nb[0] == sizeof(int8_t));
1755                return ((int8_t *)(tensor->data))[i];
1756            } break;
1757        case GGML_TYPE_I16:
1758            {
1759                GGML_ASSERT(tensor->nb[0] == sizeof(int16_t));
1760                return ((int16_t *)(tensor->data))[i];
1761            } break;
1762        case GGML_TYPE_I32:
1763            {
1764                GGML_ASSERT(tensor->nb[0] == sizeof(int32_t));
1765                return ((int32_t *)(tensor->data))[i];
1766            } break;
1767        case GGML_TYPE_F16:
1768            {
1769                GGML_ASSERT(tensor->nb[0] == sizeof(ggml_fp16_t));
1770                return GGML_FP16_TO_FP32(((ggml_fp16_t *)(tensor->data))[i]);
1771            } break;
1772        case GGML_TYPE_F32:
1773            {
1774                GGML_ASSERT(tensor->nb[0] == sizeof(float));
1775                return ((float *)(tensor->data))[i];
1776            } break;
1777        case GGML_TYPE_COUNT:
1778            {
1779                GGML_ASSERT(false);
1780            } break;
1781    }
1782
1783    return 0.0f;
1784}
1785
1786void ggml_set_i32_1d(const struct ggml_tensor * tensor, int i, int32_t value) {
1787    switch (tensor->type) {
1788        case GGML_TYPE_I8:
1789            {
1790                GGML_ASSERT(tensor->nb[0] == sizeof(int8_t));
1791                ((int8_t *)(tensor->data))[i] = value;
1792            } break;
1793        case GGML_TYPE_I16:
1794            {
1795                GGML_ASSERT(tensor->nb[0] == sizeof(int16_t));
1796                ((int16_t *)(tensor->data))[i] = value;
1797            } break;
1798        case GGML_TYPE_I32:
1799            {
1800                GGML_ASSERT(tensor->nb[0] == sizeof(int32_t));
1801                ((int32_t *)(tensor->data))[i] = value;
1802            } break;
1803        case GGML_TYPE_F16:
1804            {
1805                GGML_ASSERT(tensor->nb[0] == sizeof(ggml_fp16_t));
1806                ((ggml_fp16_t *)(tensor->data))[i] = GGML_FP32_TO_FP16(value);
1807            } break;
1808        case GGML_TYPE_F32:
1809            {
1810                GGML_ASSERT(tensor->nb[0] == sizeof(float));
1811                ((float *)(tensor->data))[i] = value;
1812            } break;
1813        case GGML_TYPE_COUNT:
1814            {
1815                GGML_ASSERT(false);
1816            } break;
1817    }
1818}
1819
1820float ggml_get_f32_1d(const struct ggml_tensor * tensor, int i) {
1821    switch (tensor->type) {
1822        case GGML_TYPE_I8:
1823            {
1824                GGML_ASSERT(tensor->nb[0] == sizeof(int8_t));
1825                return ((int8_t *)(tensor->data))[i];
1826            } break;
1827        case GGML_TYPE_I16:
1828            {
1829                GGML_ASSERT(tensor->nb[0] == sizeof(int16_t));
1830                return ((int16_t *)(tensor->data))[i];
1831            } break;
1832        case GGML_TYPE_I32:
1833            {
1834                GGML_ASSERT(tensor->nb[0] == sizeof(int32_t));
1835                return ((int32_t *)(tensor->data))[i];
1836            } break;
1837        case GGML_TYPE_F16:
1838            {
1839                GGML_ASSERT(tensor->nb[0] == sizeof(ggml_fp16_t));
1840                return GGML_FP16_TO_FP32(((ggml_fp16_t *)(tensor->data))[i]);
1841            } break;
1842        case GGML_TYPE_F32:
1843            {
1844                GGML_ASSERT(tensor->nb[0] == sizeof(float));
1845                return ((float *)(tensor->data))[i];
1846            } break;
1847        case GGML_TYPE_COUNT:
1848            {
1849                GGML_ASSERT(false);
1850            } break;
1851    }
1852
1853    return 0.0f;
1854}
1855
1856void ggml_set_f32_1d(const struct ggml_tensor * tensor, int i, float value) {
1857    switch (tensor->type) {
1858        case GGML_TYPE_I8:
1859            {
1860                GGML_ASSERT(tensor->nb[0] == sizeof(int8_t));
1861                ((int8_t *)(tensor->data))[i] = value;
1862            } break;
1863        case GGML_TYPE_I16:
1864            {
1865                GGML_ASSERT(tensor->nb[0] == sizeof(int16_t));
1866                ((int16_t *)(tensor->data))[i] = value;
1867            } break;
1868        case GGML_TYPE_I32:
1869            {
1870                GGML_ASSERT(tensor->nb[0] == sizeof(int32_t));
1871                ((int32_t *)(tensor->data))[i] = value;
1872            } break;
1873        case GGML_TYPE_F16:
1874            {
1875                GGML_ASSERT(tensor->nb[0] == sizeof(ggml_fp16_t));
1876                ((ggml_fp16_t *)(tensor->data))[i] = GGML_FP32_TO_FP16(value);
1877            } break;
1878        case GGML_TYPE_F32:
1879            {
1880                GGML_ASSERT(tensor->nb[0] == sizeof(float));
1881                ((float *)(tensor->data))[i] = value;
1882            } break;
1883        case GGML_TYPE_COUNT:
1884            {
1885                GGML_ASSERT(false);
1886            } break;
1887    }
1888}
1889
1890void * ggml_get_data(const struct ggml_tensor * tensor) {
1891    return tensor->data;
1892}
1893
1894float * ggml_get_data_f32(const struct ggml_tensor * tensor) {
1895    assert(tensor->type == GGML_TYPE_F32);
1896    return (float *)(tensor->data);
1897}
1898
1899struct ggml_tensor * ggml_view_tensor(
1900        struct ggml_context * ctx,
1901        const struct ggml_tensor * src) {
1902    return ggml_new_tensor_impl(ctx, src->type, src->n_dims, src->ne, src->data);
1903}
1904
1905////////////////////////////////////////////////////////////////////////////////
1906
1907// ggml_dup
1908
1909struct ggml_tensor * ggml_dup_impl(
1910        struct ggml_context * ctx,
1911        struct ggml_tensor * a,
1912        bool inplace) {
1913    bool is_node = false;
1914
1915    if (!inplace && (a->grad)) {
1916        is_node = true;
1917    }
1918
1919    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
1920
1921    result->op   = GGML_OP_DUP;
1922    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
1923    result->src0 = a;
1924    result->src1 = NULL;
1925
1926    return result;
1927}
1928
1929struct ggml_tensor * ggml_dup(
1930        struct ggml_context * ctx,
1931        struct ggml_tensor * a) {
1932    return ggml_dup_impl(ctx, a, false);
1933}
1934
1935struct ggml_tensor * ggml_dup_inplace(
1936        struct ggml_context * ctx,
1937        struct ggml_tensor * a) {
1938    return ggml_dup_impl(ctx, a, true);
1939}
1940
1941// ggml_add
1942
1943struct ggml_tensor * ggml_add_impl(
1944        struct ggml_context * ctx,
1945        struct ggml_tensor * a,
1946        struct ggml_tensor * b,
1947        bool inplace) {
1948    assert(ggml_are_same_shape(a, b));
1949
1950    bool is_node = false;
1951
1952    if (!inplace && (a->grad || b->grad)) {
1953        is_node = true;
1954    }
1955
1956    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
1957
1958    result->op   = GGML_OP_ADD;
1959    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
1960    result->src0 = a;
1961    result->src1 = b;
1962
1963    return result;
1964}
1965
1966struct ggml_tensor * ggml_add(
1967        struct ggml_context * ctx,
1968        struct ggml_tensor * a,
1969        struct ggml_tensor * b) {
1970    return ggml_add_impl(ctx, a, b, false);
1971}
1972
1973struct ggml_tensor * ggml_add_inplace(
1974        struct ggml_context * ctx,
1975        struct ggml_tensor * a,
1976        struct ggml_tensor * b) {
1977    return ggml_add_impl(ctx, a, b, true);
1978}
1979
1980// ggml_sub
1981
1982struct ggml_tensor * ggml_sub_impl(
1983        struct ggml_context * ctx,
1984        struct ggml_tensor * a,
1985        struct ggml_tensor * b,
1986        bool inplace) {
1987    assert(ggml_are_same_shape(a, b));
1988
1989    bool is_node = false;
1990
1991    if (!inplace && (a->grad || b->grad)) {
1992        is_node = true;
1993    }
1994
1995    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
1996
1997    result->op   = GGML_OP_SUB;
1998    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
1999    result->src0 = a;
2000    result->src1 = b;
2001
2002    return result;
2003}
2004
2005struct ggml_tensor * ggml_sub(
2006        struct ggml_context * ctx,
2007        struct ggml_tensor * a,
2008        struct ggml_tensor * b) {
2009    return ggml_sub_impl(ctx, a, b, false);
2010}
2011
2012struct ggml_tensor * ggml_sub_inplace(
2013        struct ggml_context * ctx,
2014        struct ggml_tensor * a,
2015        struct ggml_tensor * b) {
2016    return ggml_sub_impl(ctx, a, b, true);
2017}
2018
2019// ggml_mul
2020
2021struct ggml_tensor * ggml_mul_impl(
2022        struct ggml_context * ctx,
2023        struct ggml_tensor * a,
2024        struct ggml_tensor * b,
2025        bool inplace) {
2026    assert(ggml_are_same_shape(a, b));
2027
2028    bool is_node = false;
2029
2030    if (!inplace && (a->grad || b->grad)) {
2031        is_node = true;
2032    }
2033
2034    if (inplace) {
2035        assert(is_node == false);
2036    }
2037
2038    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2039
2040    result->op   = GGML_OP_MUL;
2041    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2042    result->src0 = a;
2043    result->src1 = b;
2044
2045    return result;
2046}
2047
2048struct ggml_tensor * ggml_mul(
2049        struct ggml_context * ctx,
2050        struct ggml_tensor  * a,
2051        struct ggml_tensor  * b) {
2052    return ggml_mul_impl(ctx, a, b, false);
2053}
2054
2055struct ggml_tensor * ggml_mul_inplace(
2056        struct ggml_context * ctx,
2057        struct ggml_tensor  * a,
2058        struct ggml_tensor  * b) {
2059    return ggml_mul_impl(ctx, a, b, true);
2060}
2061
2062// ggml_div
2063
2064struct ggml_tensor * ggml_div_impl(
2065        struct ggml_context * ctx,
2066        struct ggml_tensor * a,
2067        struct ggml_tensor * b,
2068        bool inplace) {
2069    assert(ggml_are_same_shape(a, b));
2070
2071    bool is_node = false;
2072
2073    if (!inplace && (a->grad || b->grad)) {
2074        is_node = true;
2075    }
2076
2077    if (inplace) {
2078        assert(is_node == false);
2079    }
2080
2081    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2082
2083    result->op   = GGML_OP_DIV;
2084    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2085    result->src0 = a;
2086    result->src1 = b;
2087
2088    return result;
2089}
2090
2091struct ggml_tensor * ggml_div(
2092        struct ggml_context * ctx,
2093        struct ggml_tensor  * a,
2094        struct ggml_tensor  * b) {
2095    return ggml_div_impl(ctx, a, b, false);
2096}
2097
2098struct ggml_tensor * ggml_div_inplace(
2099        struct ggml_context * ctx,
2100        struct ggml_tensor  * a,
2101        struct ggml_tensor  * b) {
2102    return ggml_div_impl(ctx, a, b, true);
2103}
2104
2105// ggml_sqr
2106
2107struct ggml_tensor * ggml_sqr_impl(
2108        struct ggml_context * ctx,
2109        struct ggml_tensor * a,
2110        bool inplace) {
2111    bool is_node = false;
2112
2113    if (!inplace && (a->grad)) {
2114        is_node = true;
2115    }
2116
2117    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2118
2119    result->op   = GGML_OP_SQR;
2120    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2121    result->src0 = a;
2122    result->src1 = NULL;
2123
2124    return result;
2125}
2126
2127struct ggml_tensor * ggml_sqr(
2128        struct ggml_context * ctx,
2129        struct ggml_tensor  * a) {
2130    return ggml_sqr_impl(ctx, a, false);
2131}
2132
2133struct ggml_tensor * ggml_sqr_inplace(
2134        struct ggml_context * ctx,
2135        struct ggml_tensor  * a) {
2136    return ggml_sqr_impl(ctx, a, true);
2137}
2138
2139// ggml_sqrt
2140
2141struct ggml_tensor * ggml_sqrt_impl(
2142        struct ggml_context * ctx,
2143        struct ggml_tensor * a,
2144        bool inplace) {
2145    bool is_node = false;
2146
2147    if (!inplace && (a->grad)) {
2148        is_node = true;
2149    }
2150
2151    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2152
2153    result->op   = GGML_OP_SQRT;
2154    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2155    result->src0 = a;
2156    result->src1 = NULL;
2157
2158    return result;
2159}
2160
2161struct ggml_tensor * ggml_sqrt(
2162        struct ggml_context * ctx,
2163        struct ggml_tensor  * a) {
2164    return ggml_sqrt_impl(ctx, a, false);
2165}
2166
2167struct ggml_tensor * ggml_sqrt_inplace(
2168        struct ggml_context * ctx,
2169        struct ggml_tensor  * a) {
2170    return ggml_sqrt_impl(ctx, a, true);
2171}
2172
2173// ggml_sum
2174
2175struct ggml_tensor * ggml_sum(
2176        struct ggml_context * ctx,
2177        struct ggml_tensor * a) {
2178    bool is_node = false;
2179
2180    if (a->grad) {
2181        is_node = true;
2182    }
2183
2184    struct ggml_tensor * result = ggml_new_tensor_1d(ctx, a->type, 1);
2185
2186    result->op   = GGML_OP_SUM;
2187    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2188    result->src0 = a;
2189    result->src1 = NULL;
2190
2191    return result;
2192}
2193
2194// ggml_mean
2195
2196struct ggml_tensor * ggml_mean(
2197        struct ggml_context * ctx,
2198        struct ggml_tensor * a) {
2199    bool is_node = false;
2200
2201    if (a->grad) {
2202        assert(false); // TODO: implement
2203        is_node = true;
2204    }
2205
2206    int ne[GGML_MAX_DIMS] = { 1, a->ne[1], a->ne[2], a->ne[3] };
2207    struct ggml_tensor * result = ggml_new_tensor(ctx, GGML_TYPE_F32, a->n_dims, ne);
2208
2209    result->op   = GGML_OP_MEAN;
2210    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2211    result->src0 = a;
2212    result->src1 = NULL;
2213
2214    return result;
2215}
2216
2217// ggml_repeat
2218
2219struct ggml_tensor * ggml_repeat(
2220        struct ggml_context * ctx,
2221        struct ggml_tensor * a,
2222        struct ggml_tensor * b) {
2223    assert(ggml_can_repeat(a, b));
2224
2225    bool is_node = false;
2226
2227    if (a->grad) {
2228        is_node = true;
2229    }
2230
2231    if (ggml_are_same_shape(a, b) && !is_node) {
2232        return a;
2233    }
2234
2235    struct ggml_tensor * result = ggml_new_tensor(ctx, a->type, b->n_dims, b->ne);
2236
2237    result->op   = GGML_OP_REPEAT;
2238    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2239    result->src0 = a;
2240    result->src1 = NULL;
2241
2242    return result;
2243}
2244
2245// ggml_abs
2246
2247struct ggml_tensor * ggml_abs_impl(
2248        struct ggml_context * ctx,
2249        struct ggml_tensor * a,
2250        bool inplace) {
2251    bool is_node = false;
2252
2253    if (!inplace && (a->grad)) {
2254        is_node = true;
2255    }
2256
2257    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2258
2259    result->op   = GGML_OP_ABS;
2260    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2261    result->src0 = a;
2262    result->src1 = NULL;
2263
2264    return result;
2265}
2266
2267struct ggml_tensor * ggml_abs(
2268        struct ggml_context * ctx,
2269        struct ggml_tensor  * a) {
2270    return ggml_abs_impl(ctx, a, false);
2271}
2272
2273struct ggml_tensor * ggml_abs_inplace(
2274        struct ggml_context * ctx,
2275        struct ggml_tensor  * a) {
2276    return ggml_abs_impl(ctx, a, true);
2277}
2278
2279
2280// ggml_sgn
2281
2282struct ggml_tensor * ggml_sgn_impl(
2283        struct ggml_context * ctx,
2284        struct ggml_tensor * a,
2285        bool inplace) {
2286    bool is_node = false;
2287
2288    if (!inplace && (a->grad)) {
2289        is_node = true;
2290    }
2291
2292    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2293
2294    result->op   = GGML_OP_SGN;
2295    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2296    result->src0 = a;
2297    result->src1 = NULL;
2298
2299    return result;
2300}
2301
2302struct ggml_tensor * ggml_sgn(
2303        struct ggml_context * ctx,
2304        struct ggml_tensor  * a) {
2305    return ggml_sgn_impl(ctx, a, false);
2306}
2307
2308struct ggml_tensor * ggml_sgn_inplace(
2309        struct ggml_context * ctx,
2310        struct ggml_tensor  * a) {
2311    return ggml_sgn_impl(ctx, a, true);
2312}
2313
2314// ggml_neg
2315
2316struct ggml_tensor * ggml_neg_impl(
2317        struct ggml_context * ctx,
2318        struct ggml_tensor * a,
2319        bool inplace) {
2320    bool is_node = false;
2321
2322    if (!inplace && (a->grad)) {
2323        is_node = true;
2324    }
2325
2326    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2327
2328    result->op   = GGML_OP_NEG;
2329    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2330    result->src0 = a;
2331    result->src1 = NULL;
2332
2333    return result;
2334}
2335
2336struct ggml_tensor * ggml_neg(
2337        struct ggml_context * ctx,
2338        struct ggml_tensor  * a) {
2339    return ggml_neg_impl(ctx, a, false);
2340}
2341
2342struct ggml_tensor * ggml_neg_inplace(
2343        struct ggml_context * ctx,
2344        struct ggml_tensor  * a) {
2345    return ggml_neg_impl(ctx, a, true);
2346}
2347
2348// ggml_step
2349
2350struct ggml_tensor * ggml_step_impl(
2351        struct ggml_context * ctx,
2352        struct ggml_tensor * a,
2353        bool inplace) {
2354    bool is_node = false;
2355
2356    if (!inplace && (a->grad)) {
2357        is_node = true;
2358    }
2359
2360    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2361
2362    result->op   = GGML_OP_STEP;
2363    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2364    result->src0 = a;
2365    result->src1 = NULL;
2366
2367    return result;
2368}
2369
2370struct ggml_tensor * ggml_step(
2371        struct ggml_context * ctx,
2372        struct ggml_tensor  * a) {
2373    return ggml_step_impl(ctx, a, false);
2374}
2375
2376struct ggml_tensor * ggml_step_inplace(
2377        struct ggml_context * ctx,
2378        struct ggml_tensor  * a) {
2379    return ggml_step_impl(ctx, a, true);
2380}
2381
2382// ggml_relu
2383
2384struct ggml_tensor * ggml_relu_impl(
2385        struct ggml_context * ctx,
2386        struct ggml_tensor * a,
2387        bool inplace) {
2388    bool is_node = false;
2389
2390    if (!inplace && (a->grad)) {
2391        is_node = true;
2392    }
2393
2394    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2395
2396    result->op   = GGML_OP_RELU;
2397    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2398    result->src0 = a;
2399    result->src1 = NULL;
2400
2401    return result;
2402}
2403
2404struct ggml_tensor * ggml_relu(
2405        struct ggml_context * ctx,
2406        struct ggml_tensor  * a) {
2407    return ggml_relu_impl(ctx, a, false);
2408}
2409
2410struct ggml_tensor * ggml_relu_inplace(
2411        struct ggml_context * ctx,
2412        struct ggml_tensor  * a) {
2413    return ggml_relu_impl(ctx, a, true);
2414}
2415
2416// ggml_gelu
2417
2418struct ggml_tensor * ggml_gelu_impl(
2419        struct ggml_context * ctx,
2420        struct ggml_tensor * a,
2421        bool inplace) {
2422    bool is_node = false;
2423
2424    if (!inplace && (a->grad)) {
2425        is_node = true;
2426    }
2427
2428    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2429
2430    result->op   = GGML_OP_GELU;
2431    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2432    result->src0 = a;
2433    result->src1 = NULL;
2434
2435    return result;
2436}
2437
2438struct ggml_tensor * ggml_gelu(
2439        struct ggml_context * ctx,
2440        struct ggml_tensor  * a) {
2441    return ggml_gelu_impl(ctx, a, false);
2442}
2443
2444struct ggml_tensor * ggml_gelu_inplace(
2445        struct ggml_context * ctx,
2446        struct ggml_tensor  * a) {
2447    return ggml_gelu_impl(ctx, a, true);
2448}
2449
2450// ggml_norm
2451
2452struct ggml_tensor * ggml_norm_impl(
2453        struct ggml_context * ctx,
2454        struct ggml_tensor  * a,
2455        bool inplace) {
2456    bool is_node = false;
2457
2458    if (!inplace && (a->grad)) {
2459        assert(false); // TODO: implement backward
2460        is_node = true;
2461    }
2462
2463    struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2464
2465    result->op   = GGML_OP_NORM;
2466    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2467    result->src0 = a;
2468    result->src1 = NULL; // TODO: maybe store epsilon here?
2469
2470    return result;
2471}
2472
2473struct ggml_tensor * ggml_norm(
2474        struct ggml_context * ctx,
2475        struct ggml_tensor  * a) {
2476    return ggml_norm_impl(ctx, a, false);
2477}
2478
2479struct ggml_tensor * ggml_norm_inplace(
2480        struct ggml_context * ctx,
2481        struct ggml_tensor  * a) {
2482    return ggml_norm_impl(ctx, a, true);
2483}
2484
2485// ggml_mul_mat
2486
2487struct ggml_tensor * ggml_mul_mat(
2488        struct ggml_context * ctx,
2489        struct ggml_tensor  * a,
2490        struct ggml_tensor  * b) {
2491    assert(ggml_can_mul_mat(a, b));
2492
2493	// printUniqueTensorSize( "ggml_mul_mat", a->ne, b->ne );
2494    bool is_node = false;
2495
2496    if (a->grad || b->grad) {
2497        is_node = true;
2498    }
2499
2500    const int ne[4] = { a->ne[1], b->ne[1], a->ne[2], b->ne[3] };
2501    struct ggml_tensor * result = ggml_new_tensor(ctx, GGML_TYPE_F32, MIN(a->n_dims, b->n_dims), ne);
2502
2503    result->op   = GGML_OP_MUL_MAT;
2504    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2505    result->src0 = a;
2506    result->src1 = b;
2507
2508    return result;
2509}
2510
2511// ggml_scale
2512
2513struct ggml_tensor * ggml_scale_impl(
2514        struct ggml_context * ctx,
2515        struct ggml_tensor  * a,
2516        struct ggml_tensor  * b,
2517        bool inplace) {
2518    assert(ggml_is_scalar(b));
2519    assert(ggml_is_padded_1d(a));
2520
2521    bool is_node = false;
2522
2523    if (!inplace && (a->grad || b->grad)) {
2524        assert(false); // TODO: implement backward
2525        is_node = true;
2526    }
2527
2528    // TODO: when implement backward, fix this:
2529    //struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2530    struct ggml_tensor * result = ggml_view_tensor(ctx, a);
2531
2532    result->op   = GGML_OP_SCALE;
2533    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2534    result->src0 = a;
2535    result->src1 = b;
2536
2537    return result;
2538}
2539
2540struct ggml_tensor * ggml_scale(
2541        struct ggml_context * ctx,
2542        struct ggml_tensor * a,
2543        struct ggml_tensor * b) {
2544    return ggml_scale_impl(ctx, a, b, false);
2545}
2546
2547struct ggml_tensor * ggml_scale_inplace(
2548        struct ggml_context * ctx,
2549        struct ggml_tensor * a,
2550        struct ggml_tensor * b) {
2551    return ggml_scale_impl(ctx, a, b, true);
2552}
2553
2554// ggml_cpy
2555
2556struct ggml_tensor * ggml_cpy_impl(
2557        struct ggml_context * ctx,
2558        struct ggml_tensor  * a,
2559        struct ggml_tensor  * b,
2560        bool inplace) {
2561    assert(ggml_nelements(a) == ggml_nelements(b));
2562
2563    bool is_node = false;
2564
2565    if (!inplace && (a->grad || b->grad)) {
2566        assert(false); // TODO: implement backward
2567        is_node = true;
2568    }
2569
2570    // make a view of the destination
2571    struct ggml_tensor * result = ggml_view_tensor(ctx, b);
2572
2573    result->op   = GGML_OP_CPY;
2574    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2575    result->src0 = a;
2576    result->src1 = b;
2577
2578    return result;
2579}
2580
2581struct ggml_tensor * ggml_cpy(
2582        struct ggml_context * ctx,
2583        struct ggml_tensor * a,
2584        struct ggml_tensor * b) {
2585    return ggml_cpy_impl(ctx, a, b, false);
2586}
2587
2588struct ggml_tensor * ggml_cpy_inplace(
2589        struct ggml_context * ctx,
2590        struct ggml_tensor * a,
2591        struct ggml_tensor * b) {
2592    return ggml_cpy_impl(ctx, a, b, true);
2593}
2594
2595// ggml_reshape
2596
2597struct ggml_tensor * ggml_reshape(
2598        struct ggml_context * ctx,
2599        struct ggml_tensor * a,
2600        struct ggml_tensor * b) {
2601    assert(ggml_is_contiguous(a));
2602    assert(ggml_is_contiguous(b));
2603    assert(ggml_nelements(a) == ggml_nelements(b));
2604
2605    bool is_node = false;
2606
2607    if (a->grad || b->grad) {
2608        assert(false); // TODO: implement backward
2609        is_node = true;
2610    }
2611
2612    struct ggml_tensor * result = ggml_new_tensor_impl(ctx, a->type, b->n_dims, b->ne, a->data);
2613
2614    result->op   = GGML_OP_RESHAPE;
2615    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2616    result->src0 = a;
2617    result->src1 = NULL;
2618
2619    return result;
2620}
2621
2622struct ggml_tensor * ggml_reshape_2d(
2623        struct ggml_context * ctx,
2624        struct ggml_tensor  * a,
2625        int                   ne0,
2626        int                   ne1) {
2627    assert(ggml_is_contiguous(a));
2628    assert(ggml_nelements(a) == ne0*ne1);
2629
2630    bool is_node = false;
2631
2632    if (a->grad) {
2633        assert(false); // TODO: implement backward
2634        is_node = true;
2635    }
2636
2637    const int ne[2] = { ne0, ne1 };
2638    struct ggml_tensor * result = ggml_new_tensor_impl(ctx, a->type, 2, ne, a->data);
2639
2640    result->op   = GGML_OP_RESHAPE;
2641    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2642    result->src0 = a;
2643    result->src1 = NULL;
2644
2645    return result;
2646}
2647
2648struct ggml_tensor * ggml_reshape_3d(
2649        struct ggml_context * ctx,
2650        struct ggml_tensor  * a,
2651        int                   ne0,
2652        int                   ne1,
2653        int                   ne2) {
2654    assert(ggml_is_contiguous(a));
2655    assert(ggml_nelements(a) == ne0*ne1*ne2);
2656
2657    bool is_node = false;
2658
2659    if (a->grad) {
2660        assert(false); // TODO: implement backward
2661        is_node = true;
2662    }
2663
2664    const int ne[3] = { ne0, ne1, ne2 };
2665    struct ggml_tensor * result = ggml_new_tensor_impl(ctx, a->type, 3, ne, a->data);
2666
2667    result->op   = GGML_OP_RESHAPE;
2668    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2669    result->src0 = a;
2670    result->src1 = NULL;
2671
2672    return result;
2673}
2674
2675// ggml_view_1d
2676
2677struct ggml_tensor * ggml_view_1d(
2678        struct ggml_context * ctx,
2679        struct ggml_tensor  * a,
2680        int                   ne0,
2681        size_t                offset) {
2682    if (a->grad) {
2683        assert(false); // gradient propagation is not supported
2684    }
2685
2686    struct ggml_tensor * result = ggml_new_tensor_impl(ctx, a->type, 1, &ne0, (char *) a->data + offset);
2687
2688    result->op   = GGML_OP_VIEW;
2689    result->grad = NULL;
2690    result->src0 = a;
2691    result->src1 = NULL; // TODO: maybe store the offset here?
2692
2693    return result;
2694}
2695
2696// ggml_view_2d
2697
2698struct ggml_tensor * ggml_view_2d(
2699        struct ggml_context * ctx,
2700        struct ggml_tensor  * a,
2701        int                   ne0,
2702        int                   ne1,
2703        size_t                nb1,
2704        size_t                offset) {
2705    if (a->grad) {
2706        assert(false); // gradient propagation is not supported
2707    }
2708
2709    const int ne[GGML_MAX_DIMS] = { ne0, ne1, 1, 1 };
2710
2711    struct ggml_tensor * result = ggml_new_tensor_impl(ctx, a->type, 2, ne, (char *) a->data + offset);
2712
2713    result->nb[1] = nb1;
2714    result->nb[2] = result->nb[1]*ne1;
2715    result->nb[3] = result->nb[2];
2716
2717    result->op   = GGML_OP_VIEW;
2718    result->grad = NULL;
2719    result->src0 = a;
2720    result->src1 = NULL; // TODO: maybe store the offset here?
2721
2722    return result;
2723}
2724
2725// ggml_permute
2726
2727struct ggml_tensor * ggml_permute(
2728        struct ggml_context * ctx,
2729        struct ggml_tensor  * a,
2730        int                   axis0,
2731        int                   axis1,
2732        int                   axis2,
2733        int                   axis3) {
2734    assert(axis0 >= 0 && axis0 < GGML_MAX_DIMS);
2735    assert(axis1 >= 0 && axis1 < GGML_MAX_DIMS);
2736    assert(axis2 >= 0 && axis2 < GGML_MAX_DIMS);
2737    assert(axis3 >= 0 && axis3 < GGML_MAX_DIMS);
2738
2739    assert(axis0 != axis1);
2740    assert(axis0 != axis2);
2741    assert(axis0 != axis3);
2742    assert(axis1 != axis2);
2743    assert(axis1 != axis3);
2744    assert(axis2 != axis3);
2745
2746    bool is_node = false;
2747
2748    if (a->grad) {
2749        assert(false); // TODO: implement backward
2750        is_node = true;
2751    }
2752
2753    struct ggml_tensor * result = ggml_view_tensor(ctx, a);
2754
2755    int ne[GGML_MAX_DIMS];
2756    int nb[GGML_MAX_DIMS];
2757
2758    ne[axis0] = a->ne[0];
2759    ne[axis1] = a->ne[1];
2760    ne[axis2] = a->ne[2];
2761    ne[axis3] = a->ne[3];
2762
2763    nb[axis0] = a->nb[0];
2764    nb[axis1] = a->nb[1];
2765    nb[axis2] = a->nb[2];
2766    nb[axis3] = a->nb[3];
2767
2768    result->ne[0] = ne[0];
2769    result->ne[1] = ne[1];
2770    result->ne[2] = ne[2];
2771    result->ne[3] = ne[3];
2772
2773    result->nb[0] = nb[0];
2774    result->nb[1] = nb[1];
2775    result->nb[2] = nb[2];
2776    result->nb[3] = nb[3];
2777
2778    result->op   = GGML_OP_PERMUTE;
2779    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2780    result->src0 = a;
2781    result->src1 = NULL; // TODO: maybe store the permutation here?
2782
2783    return result;
2784}
2785
2786// ggml_transpose
2787
2788struct ggml_tensor * ggml_transpose(
2789        struct ggml_context * ctx,
2790        struct ggml_tensor  * a) {
2791    bool is_node = false;
2792
2793    if (a->grad) {
2794        assert(false); // TODO: implement backward
2795        is_node = true;
2796    }
2797
2798    struct ggml_tensor * result = ggml_view_tensor(ctx, a);
2799
2800    result->ne[0] = a->ne[1];
2801    result->ne[1] = a->ne[0];
2802
2803    result->nb[0] = a->nb[1];
2804    result->nb[1] = a->nb[0];
2805
2806    result->op   = GGML_OP_TRANSPOSE;
2807    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2808    result->src0 = a;
2809    result->src1 = NULL;
2810
2811    return result;
2812}
2813
2814// ggml_get_rows
2815
2816struct ggml_tensor * ggml_get_rows(
2817        struct ggml_context * ctx,
2818        struct ggml_tensor  * a,
2819        struct ggml_tensor  * b) {
2820    assert(ggml_is_matrix(a) && ggml_is_vector(b) && b->type == GGML_TYPE_I32);
2821
2822    bool is_node = false;
2823
2824    if (a->grad || b->grad) {
2825        assert(false); // TODO: implement backward
2826        is_node = true;
2827    }
2828
2829    // TODO: implement non F32 return
2830    //struct ggml_tensor * result = ggml_new_tensor_2d(ctx, a->type, a->ne[0], b->ne[0]);
2831    struct ggml_tensor * result = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, a->ne[0], b->ne[0]);
2832
2833    result->op   = GGML_OP_GET_ROWS;
2834    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2835    result->src0 = a;
2836    result->src1 = b;
2837
2838    return result;
2839}
2840
2841// ggml_diag_mask_inf
2842
2843struct ggml_tensor * ggml_diag_mask_inf(
2844        struct ggml_context * ctx,
2845        struct ggml_tensor  * a,
2846        int                   n_past) {
2847    bool is_node = false;
2848
2849    if (a->grad) {
2850        assert(false); // TODO: implement backward
2851        is_node = true;
2852    }
2853
2854    // TODO: when implement backward, fix this:
2855    //struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2856    struct ggml_tensor * result = ggml_view_tensor(ctx, a);
2857
2858    struct ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1);
2859    ((int32_t *) b->data)[0] = n_past;
2860
2861    result->op   = GGML_OP_DIAG_MASK_INF;
2862    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2863    result->src0 = a;
2864    result->src1 = b;
2865
2866    return result;
2867}
2868
2869// ggml_soft_max
2870
2871struct ggml_tensor * ggml_soft_max(
2872        struct ggml_context * ctx,
2873        struct ggml_tensor  * a) {
2874    bool is_node = false;
2875
2876    if (a->grad) {
2877        assert(false); // TODO: implement backward
2878        is_node = true;
2879    }
2880
2881    // TODO: when implement backward, fix this:
2882    //struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2883    struct ggml_tensor * result = ggml_view_tensor(ctx, a);
2884
2885    result->op   = GGML_OP_SOFT_MAX;
2886    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2887    result->src0 = a;
2888    result->src1 = NULL;
2889
2890    return result;
2891}
2892
2893// ggml_rope
2894
2895struct ggml_tensor * ggml_rope(
2896        struct ggml_context * ctx,
2897        struct ggml_tensor  * a,
2898        int                   n_past,
2899        int                   n_dims,
2900        int                   mode) {
2901    assert(n_past >= 0);
2902    bool is_node = false;
2903
2904    if (a->grad) {
2905        assert(false); // TODO: implement backward
2906        is_node = true;
2907    }
2908
2909    // TODO: when implement backward, fix this:
2910    //struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
2911    struct ggml_tensor * result = ggml_view_tensor(ctx, a);
2912
2913    struct ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 3);
2914    ((int32_t *) b->data)[0] = n_past;
2915    ((int32_t *) b->data)[1] = n_dims;
2916    ((int32_t *) b->data)[2] = mode;
2917
2918    result->op   = GGML_OP_ROPE;
2919    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2920    result->src0 = a;
2921    result->src1 = b;
2922
2923    return result;
2924}
2925
2926// ggml_conv_1d_1s
2927
2928struct ggml_tensor * ggml_conv_1d_1s(
2929        struct ggml_context * ctx,
2930        struct ggml_tensor  * a,
2931        struct ggml_tensor  * b) {
2932    assert(ggml_is_matrix(b));
2933    assert(a->ne[1] == b->ne[1]);
2934    assert(a->ne[3] == 1);
2935    bool is_node = false;
2936
2937    if (a->grad || b->grad) {
2938        assert(false); // TODO: implement backward
2939        is_node = true;
2940    }
2941
2942    const int ne[4] = { b->ne[0], a->ne[2], 1, 1, };
2943    struct ggml_tensor * result = ggml_new_tensor(ctx, GGML_TYPE_F32, 2, ne);
2944
2945    result->op   = GGML_OP_CONV_1D_1S;
2946    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2947    result->src0 = a;
2948    result->src1 = b;
2949
2950    return result;
2951}
2952
2953// ggml_conv_1d_2s
2954
2955struct ggml_tensor * ggml_conv_1d_2s(
2956        struct ggml_context * ctx,
2957        struct ggml_tensor  * a,
2958        struct ggml_tensor  * b) {
2959    assert(ggml_is_matrix(b));
2960    assert(a->ne[1] == b->ne[1]);
2961    assert(a->ne[3] == 1);
2962    bool is_node = false;
2963
2964    if (a->grad || b->grad) {
2965        assert(false); // TODO: implement backward
2966        is_node = true;
2967    }
2968
2969    const int ne[4] = { b->ne[0]/2, a->ne[2], 1, 1, };
2970    struct ggml_tensor * result = ggml_new_tensor(ctx, GGML_TYPE_F32, 2, ne);
2971
2972    result->op   = GGML_OP_CONV_1D_2S;
2973    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
2974    result->src0 = a;
2975    result->src1 = b;
2976
2977    return result;
2978}
2979
2980// ggml_flash_attn
2981
2982struct ggml_tensor * ggml_flash_attn(
2983        struct ggml_context * ctx,
2984        struct ggml_tensor  * q,
2985        struct ggml_tensor  * k,
2986        struct ggml_tensor  * v,
2987        bool                  masked) {
2988    assert(ggml_can_mul_mat(k, q));
2989    // TODO: check if vT can be multiplied by (k*qT)
2990
2991    bool is_node = false;
2992
2993    if (q->grad || k->grad || v->grad) {
2994        GGML_ASSERT(false); // TODO: implement backward
2995        is_node = true;
2996    }
2997
2998    //struct ggml_tensor * result = ggml_dup_tensor(ctx, q);
2999    struct ggml_tensor * result = ggml_new_tensor(ctx, GGML_TYPE_F32, 4, q->ne);
3000
3001    result->op   = GGML_OP_FLASH_ATTN;
3002    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
3003    result->src0 = q;
3004    result->src1 = k;
3005    result->opt[0] = v;
3006    result->opt[1] = ggml_new_i32(ctx, masked ? 1 : 0);
3007
3008    return result;
3009}
3010
3011// ggml_flash_ff
3012
3013struct ggml_tensor * ggml_flash_ff(
3014        struct ggml_context * ctx,
3015        struct ggml_tensor  * a,
3016        struct ggml_tensor  * b0,
3017        struct ggml_tensor  * b1,
3018        struct ggml_tensor  * c0,
3019        struct ggml_tensor  * c1) {
3020    assert(ggml_can_mul_mat(b0, a));
3021    // TODO: more checks
3022
3023    bool is_node = false;
3024
3025    if (a->grad || b0->grad || b1->grad || c0->grad || c1->grad) {
3026        GGML_ASSERT(false); // TODO: implement backward
3027        is_node = true;
3028    }
3029
3030    //struct ggml_tensor * result = ggml_dup_tensor(ctx, a);
3031    struct ggml_tensor * result = ggml_new_tensor(ctx, GGML_TYPE_F32, 4, a->ne);
3032
3033    result->op   = GGML_OP_FLASH_FF;
3034    result->grad = is_node ? ggml_dup_tensor(ctx, result) : NULL;
3035    result->src0 = a;
3036    result->src1 = b0;
3037    result->opt[0] = b1;
3038    result->opt[1] = c0;
3039    result->opt[2] = c1;
3040
3041    return result;
3042}
3043
3044////////////////////////////////////////////////////////////////////////////////
3045
3046void ggml_set_param(
3047        struct ggml_context * ctx,
3048        struct ggml_tensor * tensor) {
3049    tensor->is_param = true;
3050
3051    assert(tensor->grad == NULL);
3052    tensor->grad = ggml_dup_tensor(ctx, tensor);
3053}
3054
3055// ggml_compute_forward_dup
3056
3057static void ggml_compute_forward_dup_f16(
3058        const struct ggml_compute_params * params,
3059        const struct ggml_tensor * src0,
3060        struct ggml_tensor * dst) {
3061    assert(params->ith == 0);
3062    assert(ggml_is_contiguous(dst));
3063    assert(ggml_nelements(dst) == ggml_nelements(src0));
3064
3065    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3066        return;
3067    }
3068
3069    const int ne00 = src0->ne[0];
3070    const int ne01 = src0->ne[1];
3071    const int ne02 = src0->ne[2];
3072    const int ne03 = src0->ne[3];
3073
3074    const size_t nb00 = src0->nb[0];
3075    const size_t nb01 = src0->nb[1];
3076    const size_t nb02 = src0->nb[2];
3077    const size_t nb03 = src0->nb[3];
3078
3079    if (ggml_is_contiguous(src0) && src0->type == dst->type) {
3080        memcpy(dst->data, src0->data, ggml_nelements(dst) * GGML_TYPE_SIZE[src0->type]);
3081        return;
3082    }
3083
3084    if (src0->nb[0] == sizeof(ggml_fp16_t)) {
3085        if (dst->type == GGML_TYPE_F16) {
3086            int id = 0;
3087            const size_t rs = ne00*nb00;
3088
3089            for (int i03 = 0; i03 < ne03; i03++) {
3090                for (int i02 = 0; i02 < ne02; i02++) {
3091                    for (int i01 = 0; i01 < ne01; i01++) {
3092                        const char * src0_ptr = (char *) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
3093                        char * dst_ptr = (char *) dst->data + id*rs;
3094
3095                        memcpy(dst_ptr, src0_ptr, rs);
3096
3097                        id++;
3098                    }
3099                }
3100            }
3101        } else if (dst->type == GGML_TYPE_F32) {
3102            int id = 0;
3103            float * dst_ptr = (float *) dst->data;
3104
3105            for (int i03 = 0; i03 < ne03; i03++) {
3106                for (int i02 = 0; i02 < ne02; i02++) {
3107                    for (int i01 = 0; i01 < ne01; i01++) {
3108                        for (int i00 = 0; i00 < ne00; i00++) {
3109                            const ggml_fp16_t * src0_ptr = (ggml_fp16_t *) ((char *) src0->data + i00*nb00 + i01*nb01 + i02*nb02 + i03*nb03);
3110
3111                            dst_ptr[id] = GGML_FP16_TO_FP32(*src0_ptr);
3112                            id++;
3113                        }
3114                    }
3115                }
3116            }
3117        } else {
3118            GGML_ASSERT(false); // TODO: implement
3119        }
3120    } else {
3121        //printf("%s: this is not optimal - fix me\n", __func__);
3122
3123        if (dst->type == GGML_TYPE_F32) {
3124            int id = 0;
3125            float * dst_ptr = (float *) dst->data;
3126
3127            for (int i03 = 0; i03 < ne03; i03++) {
3128                for (int i02 = 0; i02 < ne02; i02++) {
3129                    for (int i01 = 0; i01 < ne01; i01++) {
3130                        for (int i00 = 0; i00 < ne00; i00++) {
3131                            const ggml_fp16_t * src0_ptr = (ggml_fp16_t *) ((char *) src0->data + i00*nb00 + i01*nb01 + i02*nb02 + i03*nb03);
3132
3133                            dst_ptr[id] = GGML_FP16_TO_FP32(*src0_ptr);
3134                            id++;
3135                        }
3136                    }
3137                }
3138            }
3139        } else if (dst->type == GGML_TYPE_F16) {
3140            int id = 0;
3141            ggml_fp16_t * dst_ptr = (ggml_fp16_t *) dst->data;
3142
3143            for (int i03 = 0; i03 < ne03; i03++) {
3144                for (int i02 = 0; i02 < ne02; i02++) {
3145                    for (int i01 = 0; i01 < ne01; i01++) {
3146                        for (int i00 = 0; i00 < ne00; i00++) {
3147                            const ggml_fp16_t * src0_ptr = (ggml_fp16_t *) ((char *) src0->data + i00*nb00 + i01*nb01 + i02*nb02 + i03*nb03);
3148
3149                            dst_ptr[id] = *src0_ptr;
3150                            id++;
3151                        }
3152                    }
3153                }
3154            }
3155        } else {
3156            GGML_ASSERT(false); // TODO: implement
3157        }
3158    }
3159}
3160
3161static void ggml_compute_forward_dup_f32(
3162        const struct ggml_compute_params * params,
3163        const struct ggml_tensor * src0,
3164        struct ggml_tensor * dst) {
3165    GGML_ASSERT(params->ith == 0);
3166    GGML_ASSERT(ggml_is_contiguous(dst));
3167    GGML_ASSERT(ggml_nelements(dst) == ggml_nelements(src0));
3168
3169    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3170        return;
3171    }
3172
3173    const int ne00 = src0->ne[0];
3174    const int ne01 = src0->ne[1];
3175    const int ne02 = src0->ne[2];
3176    const int ne03 = src0->ne[3];
3177
3178    const size_t nb00 = src0->nb[0];
3179    const size_t nb01 = src0->nb[1];
3180    const size_t nb02 = src0->nb[2];
3181    const size_t nb03 = src0->nb[3];
3182
3183    if (ggml_is_contiguous(src0) && src0->type == dst->type) {
3184        memcpy(dst->data, src0->data, ggml_nelements(dst) * GGML_TYPE_SIZE[src0->type]);
3185        return;
3186    }
3187
3188    if (src0->nb[0] == sizeof(float)) {
3189        if (dst->type == GGML_TYPE_F32) {
3190            int id = 0;
3191            const size_t rs = ne00*nb00;
3192
3193            for (int i03 = 0; i03 < ne03; i03++) {
3194                for (int i02 = 0; i02 < ne02; i02++) {
3195                    for (int i01 = 0; i01 < ne01; i01++) {
3196                        const char * src0_ptr = (char *) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
3197                        char * dst_ptr = (char *) dst->data + id*rs;
3198
3199                        memcpy(dst_ptr, src0_ptr, rs);
3200
3201                        id++;
3202                    }
3203                }
3204            }
3205        } else if (dst->type == GGML_TYPE_F16) {
3206            int id = 0;
3207            ggml_fp16_t * dst_ptr = (ggml_fp16_t *) dst->data;
3208
3209            for (int i03 = 0; i03 < ne03; i03++) {
3210                for (int i02 = 0; i02 < ne02; i02++) {
3211                    for (int i01 = 0; i01 < ne01; i01++) {
3212                        for (int i00 = 0; i00 < ne00; i00++) {
3213                            const float * src0_ptr = (float *) ((char *) src0->data + i00*nb00 + i01*nb01 + i02*nb02 + i03*nb03);
3214
3215                            dst_ptr[id] = GGML_FP32_TO_FP16(*src0_ptr);
3216                            id++;
3217                        }
3218                    }
3219                }
3220            }
3221        } else {
3222            GGML_ASSERT(false); // TODO: implement
3223        }
3224    } else {
3225        //printf("%s: this is not optimal - fix me\n", __func__);
3226
3227        if (dst->type == GGML_TYPE_F32) {
3228            int id = 0;
3229            float * dst_ptr = (float *) dst->data;
3230
3231            for (int i03 = 0; i03 < ne03; i03++) {
3232                for (int i02 = 0; i02 < ne02; i02++) {
3233                    for (int i01 = 0; i01 < ne01; i01++) {
3234                        for (int i00 = 0; i00 < ne00; i00++) {
3235                            const float * src0_ptr = (float *) ((char *) src0->data + i00*nb00 + i01*nb01 + i02*nb02 + i03*nb03);
3236
3237                            dst_ptr[id] = *src0_ptr;
3238                            id++;
3239                        }
3240                    }
3241                }
3242            }
3243        } else if (dst->type == GGML_TYPE_F16) {
3244            int id = 0;
3245            ggml_fp16_t * dst_ptr = (ggml_fp16_t *) dst->data;
3246
3247            for (int i03 = 0; i03 < ne03; i03++) {
3248                for (int i02 = 0; i02 < ne02; i02++) {
3249                    for (int i01 = 0; i01 < ne01; i01++) {
3250                        for (int i00 = 0; i00 < ne00; i00++) {
3251                            const float * src0_ptr = (float *) ((char *) src0->data + i00*nb00 + i01*nb01 + i02*nb02 + i03*nb03);
3252
3253                            dst_ptr[id] = GGML_FP32_TO_FP16(*src0_ptr);
3254                            id++;
3255                        }
3256                    }
3257                }
3258            }
3259        } else {
3260            GGML_ASSERT(false); // TODO: implement
3261        }
3262    }
3263}
3264
3265static void ggml_compute_forward_dup(
3266        const struct ggml_compute_params * params,
3267        const struct ggml_tensor * src0,
3268        struct ggml_tensor * dst) {
3269    switch (src0->type) {
3270        case GGML_TYPE_F16:
3271            {
3272                ggml_compute_forward_dup_f16(params, src0, dst);
3273            } break;
3274        case GGML_TYPE_F32:
3275            {
3276                ggml_compute_forward_dup_f32(params, src0, dst);
3277            } break;
3278        case GGML_TYPE_I8:
3279        case GGML_TYPE_I16:
3280        case GGML_TYPE_I32:
3281        case GGML_TYPE_COUNT:
3282            {
3283                GGML_ASSERT(false);
3284            } break;
3285    }
3286}
3287
3288// ggml_compute_forward_add
3289
3290static void ggml_compute_forward_add_f32(
3291        const struct ggml_compute_params * params,
3292        const struct ggml_tensor * src0,
3293        const struct ggml_tensor * src1,
3294        struct ggml_tensor * dst) {
3295    GGML_ASSERT(ggml_are_same_shape(src0, src1) && ggml_are_same_shape(src0, dst));
3296
3297    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3298        return;
3299    }
3300
3301    const int ith = params->ith;
3302    const int nth = params->nth;
3303
3304    const int n  = ggml_nrows(src0);
3305    const int nc = src0->ne[0];
3306
3307    const size_t nb00 = src0->nb[0];
3308    const size_t nb01 = src0->nb[1];
3309
3310    const size_t nb10 = src1->nb[0];
3311    const size_t nb11 = src1->nb[1];
3312
3313    const size_t nb0 = dst->nb[0];
3314    const size_t nb1 = dst->nb[1];
3315
3316    GGML_ASSERT( nb0 == sizeof(float));
3317    GGML_ASSERT(nb00 == sizeof(float));
3318
3319    if (nb10 == sizeof(float)) {
3320        const int j0 = (n/nth)*ith;
3321        const int j1 = ith == nth - 1 ? n : (n/nth)*(ith + 1);
3322
3323        for (int j = j0; j < j1; j++) {
3324            ggml_vec_add_f32(nc,
3325                    (float *) ((char *) dst->data  + j*nb1),
3326                    (float *) ((char *) src0->data + j*nb01),
3327                    (float *) ((char *) src1->data + j*nb11));
3328        }
3329    } else {
3330        // src1 is not contiguous
3331        for (int j = ith; j < n; j += nth) {
3332            float * dst_ptr  = (float *) ((char *) dst->data  + j*nb1);
3333            float * src0_ptr = (float *) ((char *) src0->data + j*nb01);
3334            for (int i = 0; i < nc; i++) {
3335                float * src1_ptr = (float *) ((char *) src1->data + j*nb11 + i*nb10);
3336
3337                dst_ptr[i] = src0_ptr[i] + *src1_ptr;
3338            }
3339        }
3340    }
3341}
3342
3343static void ggml_compute_forward_add(
3344        const struct ggml_compute_params * params,
3345        const struct ggml_tensor * src0,
3346        const struct ggml_tensor * src1,
3347        struct ggml_tensor * dst) {
3348    switch (src0->type) {
3349        case GGML_TYPE_F32:
3350            {
3351                ggml_compute_forward_add_f32(params, src0, src1, dst);
3352            } break;
3353        case GGML_TYPE_I8:
3354        case GGML_TYPE_I16:
3355        case GGML_TYPE_I32:
3356        case GGML_TYPE_F16:
3357        case GGML_TYPE_COUNT:
3358            {
3359                assert(false);
3360            } break;
3361    }
3362}
3363
3364// ggml_compute_forward_sub
3365
3366static void ggml_compute_forward_sub_f32(
3367        const struct ggml_compute_params * params,
3368        const struct ggml_tensor * src0,
3369        const struct ggml_tensor * src1,
3370        struct ggml_tensor * dst) {
3371    assert(params->ith == 0);
3372    assert(ggml_are_same_shape(src0, src1) && ggml_are_same_shape(src0, dst));
3373
3374    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3375        return;
3376    }
3377
3378    const int n  = ggml_nrows(src0);
3379    const int nc = src0->ne[0];
3380
3381    assert( dst->nb[0] == sizeof(float));
3382    assert(src0->nb[0] == sizeof(float));
3383    assert(src1->nb[0] == sizeof(float));
3384
3385    for (int i = 0; i < n; i++) {
3386        ggml_vec_sub_f32(nc,
3387                (float *) ((char *) dst->data  + i*( dst->nb[1])),
3388                (float *) ((char *) src0->data + i*(src0->nb[1])),
3389                (float *) ((char *) src1->data + i*(src1->nb[1])));
3390    }
3391}
3392
3393static void ggml_compute_forward_sub(
3394        const struct ggml_compute_params * params,
3395        const struct ggml_tensor * src0,
3396        const struct ggml_tensor * src1,
3397        struct ggml_tensor * dst) {
3398    switch (src0->type) {
3399        case GGML_TYPE_F32:
3400            {
3401                ggml_compute_forward_sub_f32(params, src0, src1, dst);
3402            } break;
3403        case GGML_TYPE_I8:
3404        case GGML_TYPE_I16:
3405        case GGML_TYPE_I32:
3406        case GGML_TYPE_F16:
3407        case GGML_TYPE_COUNT:
3408            {
3409                assert(false);
3410            } break;
3411    }
3412}
3413
3414// ggml_compute_forward_mul
3415
3416static void ggml_compute_forward_mul_f32(
3417        const struct ggml_compute_params * params,
3418        const struct ggml_tensor * src0,
3419        const struct ggml_tensor * src1,
3420        struct ggml_tensor * dst) {
3421    assert(params->ith == 0);
3422    assert(ggml_are_same_shape(src0, src1) && ggml_are_same_shape(src0, dst));
3423
3424    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3425        return;
3426    }
3427
3428    const int n  = ggml_nrows(src0);
3429    const int nc = src0->ne[0];
3430
3431    assert( dst->nb[0] == sizeof(float));
3432    assert(src0->nb[0] == sizeof(float));
3433    assert(src1->nb[0] == sizeof(float));
3434
3435    for (int i = 0; i < n; i++) {
3436        ggml_vec_mul_f32(nc,
3437                (float *) ((char *) dst->data  + i*( dst->nb[1])),
3438                (float *) ((char *) src0->data + i*(src0->nb[1])),
3439                (float *) ((char *) src1->data + i*(src1->nb[1])));
3440    }
3441}
3442
3443static void ggml_compute_forward_mul(
3444        const struct ggml_compute_params * params,
3445        const struct ggml_tensor * src0,
3446        const struct ggml_tensor * src1,
3447        struct ggml_tensor * dst) {
3448    switch (src0->type) {
3449        case GGML_TYPE_F32:
3450            {
3451                ggml_compute_forward_mul_f32(params, src0, src1, dst);
3452            } break;
3453        case GGML_TYPE_I8:
3454        case GGML_TYPE_I16:
3455        case GGML_TYPE_I32:
3456        case GGML_TYPE_F16:
3457        case GGML_TYPE_COUNT:
3458            {
3459                assert(false);
3460            } break;
3461    }
3462}
3463
3464// ggml_compute_forward_div
3465
3466static void ggml_compute_forward_div_f32(
3467        const struct ggml_compute_params * params,
3468        const struct ggml_tensor * src0,
3469        const struct ggml_tensor * src1,
3470        struct ggml_tensor * dst) {
3471    assert(params->ith == 0);
3472    assert(ggml_are_same_shape(src0, src1) && ggml_are_same_shape(src0, dst));
3473
3474    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3475        return;
3476    }
3477
3478    const int n  = ggml_nrows(src0);
3479    const int nc = src0->ne[0];
3480
3481    assert( dst->nb[0] == sizeof(float));
3482    assert(src0->nb[0] == sizeof(float));
3483    assert(src1->nb[0] == sizeof(float));
3484
3485    for (int i = 0; i < n; i++) {
3486        ggml_vec_div_f32(nc,
3487                (float *) ((char *) dst->data  + i*( dst->nb[1])),
3488                (float *) ((char *) src0->data + i*(src0->nb[1])),
3489                (float *) ((char *) src1->data + i*(src1->nb[1])));
3490    }
3491}
3492
3493static void ggml_compute_forward_div(
3494        const struct ggml_compute_params * params,
3495        const struct ggml_tensor * src0,
3496        const struct ggml_tensor * src1,
3497        struct ggml_tensor * dst) {
3498    switch (src0->type) {
3499        case GGML_TYPE_F32:
3500            {
3501                ggml_compute_forward_div_f32(params, src0, src1, dst);
3502            } break;
3503        case GGML_TYPE_I8:
3504        case GGML_TYPE_I16:
3505        case GGML_TYPE_I32:
3506        case GGML_TYPE_F16:
3507        case GGML_TYPE_COUNT:
3508            {
3509                assert(false);
3510            } break;
3511    }
3512}
3513
3514// ggml_compute_forward_sqr
3515
3516static void ggml_compute_forward_sqr_f32(
3517        const struct ggml_compute_params * params,
3518        const struct ggml_tensor * src0,
3519        struct ggml_tensor * dst) {
3520    assert(params->ith == 0);
3521    assert(ggml_are_same_shape(src0, dst));
3522
3523    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3524        return;
3525    }
3526
3527    const int n     = ggml_nrows(src0);
3528    const int nc    = src0->ne[0];
3529
3530    assert( dst->nb[0] == sizeof(float));
3531    assert(src0->nb[0] == sizeof(float));
3532
3533    for (int i = 0; i < n; i++) {
3534        ggml_vec_sqr_f32(nc,
3535                (float *) ((char *) dst->data  + i*( dst->nb[1])),
3536                (float *) ((char *) src0->data + i*(src0->nb[1])));
3537    }
3538}
3539
3540static void ggml_compute_forward_sqr(
3541        const struct ggml_compute_params * params,
3542        const struct ggml_tensor * src0,
3543        struct ggml_tensor * dst) {
3544    switch (src0->type) {
3545        case GGML_TYPE_F32:
3546            {
3547                ggml_compute_forward_sqr_f32(params, src0, dst);
3548            } break;
3549        case GGML_TYPE_I8:
3550        case GGML_TYPE_I16:
3551        case GGML_TYPE_I32:
3552        case GGML_TYPE_F16:
3553        case GGML_TYPE_COUNT:
3554            {
3555                assert(false);
3556            } break;
3557    }
3558}
3559
3560// ggml_compute_forward_sqrt
3561
3562static void ggml_compute_forward_sqrt_f32(
3563        const struct ggml_compute_params * params,
3564        const struct ggml_tensor * src0,
3565        struct ggml_tensor * dst) {
3566    assert(params->ith == 0);
3567    assert(ggml_are_same_shape(src0, dst));
3568
3569    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3570        return;
3571    }
3572
3573    const int n  = ggml_nrows(src0);
3574    const int nc = src0->ne[0];
3575
3576    assert( dst->nb[0] == sizeof(float));
3577    assert(src0->nb[0] == sizeof(float));
3578
3579    for (int i = 0; i < n; i++) {
3580        ggml_vec_sqrt_f32(nc,
3581                (float *) ((char *) dst->data  + i*( dst->nb[1])),
3582                (float *) ((char *) src0->data + i*(src0->nb[1])));
3583    }
3584}
3585
3586static void ggml_compute_forward_sqrt(
3587        const struct ggml_compute_params * params,
3588        const struct ggml_tensor * src0,
3589        struct ggml_tensor * dst) {
3590    switch (src0->type) {
3591        case GGML_TYPE_F32:
3592            {
3593                ggml_compute_forward_sqrt_f32(params, src0, dst);
3594            } break;
3595        case GGML_TYPE_I8:
3596        case GGML_TYPE_I16:
3597        case GGML_TYPE_I32:
3598        case GGML_TYPE_F16:
3599        case GGML_TYPE_COUNT:
3600            {
3601                assert(false);
3602            } break;
3603    }
3604}
3605
3606// ggml_compute_forward_sum
3607
3608static void ggml_compute_forward_sum_f32(
3609        const struct ggml_compute_params * params,
3610        const struct ggml_tensor * src0,
3611        struct ggml_tensor * dst) {
3612    assert(params->ith == 0);
3613    assert(ggml_is_scalar(dst));
3614
3615    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3616        return;
3617    }
3618
3619    assert(ggml_is_scalar(dst));
3620    assert(src0->nb[0] == sizeof(float));
3621
3622    *(float *) (dst->data) = 0.0f;
3623
3624    const int ne00 = src0->ne[0];
3625    const int ne01 = src0->ne[1];
3626    const int ne02 = src0->ne[2];
3627    const int ne03 = src0->ne[3];
3628
3629    const size_t nb01 = src0->nb[1];
3630    const size_t nb02 = src0->nb[2];
3631    const size_t nb03 = src0->nb[3];
3632
3633    for (int i03 = 0; i03 < ne03; i03++) {
3634        for (int i02 = 0; i02 < ne02; i02++) {
3635            for (int i01 = 0; i01 < ne01; i01++) {
3636                ggml_vec_sum_f32(ne00,
3637                        (float *) (dst->data),
3638                        (float *) ((char *) src0->data + i01*nb01 + i02*nb02 + i03*nb03));
3639            }
3640        }
3641    }
3642}
3643
3644static void ggml_compute_forward_sum(
3645        const struct ggml_compute_params * params,
3646        const struct ggml_tensor * src0,
3647        struct ggml_tensor * dst) {
3648    switch (src0->type) {
3649        case GGML_TYPE_F32:
3650            {
3651                ggml_compute_forward_sum_f32(params, src0, dst);
3652            } break;
3653        case GGML_TYPE_I8:
3654        case GGML_TYPE_I16:
3655        case GGML_TYPE_I32:
3656        case GGML_TYPE_F16:
3657        case GGML_TYPE_COUNT:
3658            {
3659                assert(false);
3660            } break;
3661    }
3662}
3663
3664// ggml_compute_forward_mean
3665
3666static void ggml_compute_forward_mean_f32(
3667        const struct ggml_compute_params * params,
3668        const struct ggml_tensor * src0,
3669        struct ggml_tensor * dst) {
3670    assert(params->ith == 0);
3671
3672    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3673        return;
3674    }
3675
3676    assert(src0->nb[0] == sizeof(float));
3677
3678    const int ne00 = src0->ne[0];
3679    const int ne01 = src0->ne[1];
3680    const int ne02 = src0->ne[2];
3681    const int ne03 = src0->ne[3];
3682
3683    const size_t nb01 = src0->nb[1];
3684    const size_t nb02 = src0->nb[2];
3685    const size_t nb03 = src0->nb[3];
3686
3687    const int ne0 = dst->ne[0];
3688    const int ne1 = dst->ne[1];
3689    const int ne2 = dst->ne[2];
3690    const int ne3 = dst->ne[3];
3691
3692    assert(ne0 == 1);
3693    assert(ne1 == ne01);
3694    assert(ne2 == ne02);
3695    assert(ne3 == ne03);
3696
3697    UNUSED(ne0);
3698    UNUSED(ne1);
3699    UNUSED(ne2);
3700    UNUSED(ne3);
3701
3702    const size_t nb1 = dst->nb[1];
3703    const size_t nb2 = dst->nb[2];
3704    const size_t nb3 = dst->nb[3];
3705
3706    for (int i03 = 0; i03 < ne03; i03++) {
3707        for (int i02 = 0; i02 < ne02; i02++) {
3708            for (int i01 = 0; i01 < ne01; i01++) {
3709                *(float *) ((char *) dst->data + i01*nb1 + i02*nb2 + i03*nb3) = 0.0f;
3710
3711                ggml_vec_sum_f32(ne00,
3712                        (float *) ((char *)  dst->data + i01*nb1  + i02*nb2  + i03*nb3),
3713                        (float *) ((char *) src0->data + i01*nb01 + i02*nb02 + i03*nb03));
3714
3715                *(float *) ((char *) dst->data + i01*nb1 + i02*nb2 + i03*nb3) /= (float) ne00;
3716            }
3717        }
3718    }
3719}
3720
3721static void ggml_compute_forward_mean(
3722        const struct ggml_compute_params * params,
3723        const struct ggml_tensor * src0,
3724        struct ggml_tensor * dst) {
3725    switch (src0->type) {
3726        case GGML_TYPE_F32:
3727            {
3728                ggml_compute_forward_mean_f32(params, src0, dst);
3729            } break;
3730        case GGML_TYPE_I8:
3731        case GGML_TYPE_I16:
3732        case GGML_TYPE_I32:
3733        case GGML_TYPE_F16:
3734        case GGML_TYPE_COUNT:
3735            {
3736                assert(false);
3737            } break;
3738    }
3739}
3740
3741// ggml_compute_forward_repeat
3742
3743static void ggml_compute_forward_repeat_f32(
3744        const struct ggml_compute_params * params,
3745        const struct ggml_tensor * src0,
3746        struct ggml_tensor * dst) {
3747    assert(params->ith == 0);
3748    assert(ggml_can_repeat(src0, dst));
3749
3750    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3751        return;
3752    }
3753
3754    // TODO: implement support for rank > 2 tensors
3755    assert(src0->ne[2] == 1);
3756    assert(src0->ne[3] == 1);
3757    assert( dst->ne[2] == 1);
3758    assert( dst->ne[3] == 1);
3759
3760    const int nc  = dst->ne[0];
3761    const int nr  = dst->ne[1];
3762    const int nc0 = src0->ne[0];
3763    const int nr0 = src0->ne[1];
3764    const int ncr = nc/nc0; // guaranteed to be an integer due to the check in ggml_can_repeat
3765    const int nrr = nr/nr0; // guaranteed to be an integer due to the check in ggml_can_repeat
3766
3767    // TODO: support for transposed / permuted tensors
3768    assert( dst->nb[0] == sizeof(float));
3769    assert(src0->nb[0] == sizeof(float));
3770
3771    // TODO: maybe this is not optimal?
3772    for (int i = 0; i < nrr; i++) {
3773        for (int j = 0; j < ncr; j++) {
3774            for (int k = 0; k < nr0; k++) {
3775                ggml_vec_cpy_f32(nc0,
3776                        (float *) ((char *)  dst->data + (i*nr0 + k)*( dst->nb[1]) + j*nc0*( dst->nb[0])),
3777                        (float *) ((char *) src0->data + (        k)*(src0->nb[1])));
3778            }
3779        }
3780    }
3781}
3782
3783static void ggml_compute_forward_repeat(
3784        const struct ggml_compute_params * params,
3785        const struct ggml_tensor * src0,
3786        struct ggml_tensor * dst) {
3787    switch (src0->type) {
3788        case GGML_TYPE_F32:
3789            {
3790                ggml_compute_forward_repeat_f32(params, src0, dst);
3791            } break;
3792        case GGML_TYPE_I8:
3793        case GGML_TYPE_I16:
3794        case GGML_TYPE_I32:
3795        case GGML_TYPE_F16:
3796        case GGML_TYPE_COUNT:
3797            {
3798                assert(false);
3799            } break;
3800    }
3801}
3802
3803// ggml_compute_forward_abs
3804
3805static void ggml_compute_forward_abs_f32(
3806        const struct ggml_compute_params * params,
3807        const struct ggml_tensor * src0,
3808        struct ggml_tensor * dst) {
3809    assert(params->ith == 0);
3810    assert(ggml_are_same_shape(src0, dst));
3811
3812    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3813        return;
3814    }
3815
3816    const int n  = ggml_nrows(src0);
3817    const int nc = src0->ne[0];
3818
3819    assert(dst->nb[0]  == sizeof(float));
3820    assert(src0->nb[0] == sizeof(float));
3821
3822    for (int i = 0; i < n; i++) {
3823        ggml_vec_abs_f32(nc,
3824                (float *) ((char *) dst->data  + i*( dst->nb[1])),
3825                (float *) ((char *) src0->data + i*(src0->nb[1])));
3826    }
3827}
3828
3829static void ggml_compute_forward_abs(
3830        const struct ggml_compute_params * params,
3831        const struct ggml_tensor * src0,
3832        struct ggml_tensor * dst) {
3833    switch (src0->type) {
3834        case GGML_TYPE_F32:
3835            {
3836                ggml_compute_forward_abs_f32(params, src0, dst);
3837            } break;
3838        case GGML_TYPE_I8:
3839        case GGML_TYPE_I16:
3840        case GGML_TYPE_I32:
3841        case GGML_TYPE_F16:
3842        case GGML_TYPE_COUNT:
3843            {
3844                assert(false);
3845            } break;
3846    }
3847}
3848
3849// ggml_compute_forward_sgn
3850
3851static void ggml_compute_forward_sgn_f32(
3852        const struct ggml_compute_params * params,
3853        const struct ggml_tensor * src0,
3854        struct ggml_tensor * dst) {
3855    assert(params->ith == 0);
3856    assert(ggml_are_same_shape(src0, dst));
3857
3858    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3859        return;
3860    }
3861
3862    const int n  = ggml_nrows(src0);
3863    const int nc = src0->ne[0];
3864
3865    assert(dst->nb[0]  == sizeof(float));
3866    assert(src0->nb[0] == sizeof(float));
3867
3868    for (int i = 0; i < n; i++) {
3869        ggml_vec_sgn_f32(nc,
3870                (float *) ((char *) dst->data  + i*( dst->nb[1])),
3871                (float *) ((char *) src0->data + i*(src0->nb[1])));
3872    }
3873}
3874
3875static void ggml_compute_forward_sgn(
3876        const struct ggml_compute_params * params,
3877        const struct ggml_tensor * src0,
3878        struct ggml_tensor * dst) {
3879    switch (src0->type) {
3880        case GGML_TYPE_F32:
3881            {
3882                ggml_compute_forward_sgn_f32(params, src0, dst);
3883            } break;
3884        case GGML_TYPE_I8:
3885        case GGML_TYPE_I16:
3886        case GGML_TYPE_I32:
3887        case GGML_TYPE_F16:
3888        case GGML_TYPE_COUNT:
3889            {
3890                assert(false);
3891            } break;
3892    }
3893}
3894
3895// ggml_compute_forward_neg
3896
3897static void ggml_compute_forward_neg_f32(
3898        const struct ggml_compute_params * params,
3899        const struct ggml_tensor * src0,
3900        struct ggml_tensor * dst) {
3901    assert(params->ith == 0);
3902    assert(ggml_are_same_shape(src0, dst));
3903
3904    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3905        return;
3906    }
3907
3908    const int n  = ggml_nrows(src0);
3909    const int nc = src0->ne[0];
3910
3911    assert(dst->nb[0]  == sizeof(float));
3912    assert(src0->nb[0] == sizeof(float));
3913
3914    for (int i = 0; i < n; i++) {
3915        ggml_vec_neg_f32(nc,
3916                (float *) ((char *) dst->data  + i*( dst->nb[1])),
3917                (float *) ((char *) src0->data + i*(src0->nb[1])));
3918    }
3919}
3920
3921static void ggml_compute_forward_neg(
3922        const struct ggml_compute_params * params,
3923        const struct ggml_tensor * src0,
3924        struct ggml_tensor * dst) {
3925    switch (src0->type) {
3926        case GGML_TYPE_F32:
3927            {
3928                ggml_compute_forward_neg_f32(params, src0, dst);
3929            } break;
3930        case GGML_TYPE_I8:
3931        case GGML_TYPE_I16:
3932        case GGML_TYPE_I32:
3933        case GGML_TYPE_F16:
3934        case GGML_TYPE_COUNT:
3935            {
3936                assert(false);
3937            } break;
3938    }
3939}
3940
3941// ggml_compute_forward_step
3942
3943static void ggml_compute_forward_step_f32(
3944        const struct ggml_compute_params * params,
3945        const struct ggml_tensor * src0,
3946        struct ggml_tensor * dst) {
3947    assert(params->ith == 0);
3948    assert(ggml_are_same_shape(src0, dst));
3949
3950    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3951        return;
3952    }
3953
3954    const int n  = ggml_nrows(src0);
3955    const int nc = src0->ne[0];
3956
3957    assert(dst->nb[0]  == sizeof(float));
3958    assert(src0->nb[0] == sizeof(float));
3959
3960    for (int i = 0; i < n; i++) {
3961        ggml_vec_step_f32(nc,
3962                (float *) ((char *) dst->data  + i*( dst->nb[1])),
3963                (float *) ((char *) src0->data + i*(src0->nb[1])));
3964    }
3965}
3966
3967static void ggml_compute_forward_step(
3968        const struct ggml_compute_params * params,
3969        const struct ggml_tensor * src0,
3970        struct ggml_tensor * dst) {
3971    switch (src0->type) {
3972        case GGML_TYPE_F32:
3973            {
3974                ggml_compute_forward_step_f32(params, src0, dst);
3975            } break;
3976        case GGML_TYPE_I8:
3977        case GGML_TYPE_I16:
3978        case GGML_TYPE_I32:
3979        case GGML_TYPE_F16:
3980        case GGML_TYPE_COUNT:
3981            {
3982                assert(false);
3983            } break;
3984    }
3985}
3986
3987// ggml_compute_forward_relu
3988
3989static void ggml_compute_forward_relu_f32(
3990        const struct ggml_compute_params * params,
3991        const struct ggml_tensor * src0,
3992        struct ggml_tensor * dst) {
3993    assert(params->ith == 0);
3994    assert(ggml_are_same_shape(src0, dst));
3995
3996    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
3997        return;
3998    }
3999
4000    const int n  = ggml_nrows(src0);
4001    const int nc = src0->ne[0];
4002
4003    assert(dst->nb[0]  == sizeof(float));
4004    assert(src0->nb[0] == sizeof(float));
4005
4006    for (int i = 0; i < n; i++) {
4007        ggml_vec_relu_f32(nc,
4008                (float *) ((char *) dst->data  + i*( dst->nb[1])),
4009                (float *) ((char *) src0->data + i*(src0->nb[1])));
4010    }
4011}
4012
4013static void ggml_compute_forward_relu(
4014        const struct ggml_compute_params * params,
4015        const struct ggml_tensor * src0,
4016        struct ggml_tensor * dst) {
4017    switch (src0->type) {
4018        case GGML_TYPE_F32:
4019            {
4020                ggml_compute_forward_relu_f32(params, src0, dst);
4021            } break;
4022        case GGML_TYPE_I8:
4023        case GGML_TYPE_I16:
4024        case GGML_TYPE_I32:
4025        case GGML_TYPE_F16:
4026        case GGML_TYPE_COUNT:
4027            {
4028                assert(false);
4029            } break;
4030    }
4031}
4032
4033// ggml_compute_forward_gelu
4034
4035static void ggml_compute_forward_gelu_f32(
4036        const struct ggml_compute_params * params,
4037        const struct ggml_tensor * src0,
4038        struct ggml_tensor * dst) {
4039    GGML_ASSERT(ggml_is_contiguous(src0));
4040    GGML_ASSERT(ggml_is_contiguous(dst));
4041    GGML_ASSERT(ggml_are_same_shape(src0, dst));
4042
4043    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
4044        return;
4045    }
4046
4047    const int ith = params->ith;
4048    const int nth = params->nth;
4049
4050    const int nc = src0->ne[0];
4051    const int nr = ggml_nrows(src0);
4052
4053    // rows per thread
4054    const int dr = (nr + nth - 1)/nth;
4055
4056    // row range for this thread
4057    const int ir0 = dr*ith;
4058    const int ir1 = MIN(ir0 + dr, nr);
4059
4060    for (int i1 = ir0; i1 < ir1; i1++) {
4061        ggml_vec_gelu_f32(nc,
4062                (float *) ((char *) dst->data  + i1*( dst->nb[1])),
4063                (float *) ((char *) src0->data + i1*(src0->nb[1])));
4064
4065#ifndef NDEBUG
4066        for (int k = 0; k < nc; k++) {
4067            const float x = ((float *) ((char *) dst->data + i1*( dst->nb[1])))[k];
4068            UNUSED(x);
4069            assert(!isnan(x));
4070            assert(!isinf(x));
4071        }
4072#endif
4073    }
4074}
4075
4076static void ggml_compute_forward_gelu(
4077        const struct ggml_compute_params * params,
4078        const struct ggml_tensor * src0,
4079        struct ggml_tensor * dst) {
4080    switch (src0->type) {
4081        case GGML_TYPE_F32:
4082            {
4083                ggml_compute_forward_gelu_f32(params, src0, dst);
4084            } break;
4085        case GGML_TYPE_I8:
4086        case GGML_TYPE_I16:
4087        case GGML_TYPE_I32:
4088        case GGML_TYPE_F16:
4089        case GGML_TYPE_COUNT:
4090            {
4091                assert(false);
4092            } break;
4093    }
4094}
4095
4096// ggml_compute_forward_norm
4097
4098static void ggml_compute_forward_norm_f32(
4099        const struct ggml_compute_params * params,
4100        const struct ggml_tensor * src0,
4101        struct ggml_tensor * dst) {
4102    GGML_ASSERT(ggml_are_same_shape(src0, dst));
4103
4104    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
4105        return;
4106    }
4107
4108    GGML_ASSERT(src0->nb[0] == sizeof(float));
4109
4110    const int ith = params->ith;
4111    const int nth = params->nth;
4112
4113    const int ne00 = src0->ne[0];
4114    const int ne01 = src0->ne[1];
4115    const int ne02 = src0->ne[2];
4116    const int ne03 = src0->ne[3];
4117
4118    const size_t nb01 = src0->nb[1];
4119    const size_t nb02 = src0->nb[2];
4120    const size_t nb03 = src0->nb[3];
4121
4122    const size_t nb1 = dst->nb[1];
4123    const size_t nb2 = dst->nb[2];
4124    const size_t nb3 = dst->nb[3];
4125
4126    const ggml_float eps = 1e-5f; // TODO: make this a parameter
4127
4128    // TODO: optimize
4129    for (int i03 = 0; i03 < ne03; i03++) {
4130        for (int i02 = 0; i02 < ne02; i02++) {
4131            for (int i01 = ith; i01 < ne01; i01 += nth) {
4132                const float * x = (float *) ((char *) src0->data + i01*nb01 + i02*nb02 + i03*nb03);
4133
4134                ggml_float mean = 0.0;
4135                for (int i00 = 0; i00 < ne00; i00++) {
4136                    mean += x[i00];
4137                }
4138
4139                mean /= ne00;
4140
4141                float * y = (float *) ((char *) dst->data + i01*nb1 + i02*nb2 + i03*nb3);
4142
4143                ggml_float sum2 = 0.0;
4144                for (int i00 = 0; i00 < ne00; i00++) {
4145                    ggml_float v = x[i00] - mean;
4146                    y[i00] = v;
4147                    sum2 += v*v;
4148                }
4149
4150                const float scale = 1.0/sqrt(sum2/ne00 + eps);
4151
4152                ggml_vec_scale_f32(ne00, y, scale);
4153            }
4154        }
4155    }
4156}
4157
4158static void ggml_compute_forward_norm(
4159        const struct ggml_compute_params * params,
4160        const struct ggml_tensor * src0,
4161        struct ggml_tensor * dst) {
4162    switch (src0->type) {
4163        case GGML_TYPE_F32:
4164            {
4165                ggml_compute_forward_norm_f32(params, src0, dst);
4166            } break;
4167        case GGML_TYPE_I8:
4168        case GGML_TYPE_I16:
4169        case GGML_TYPE_I32:
4170        case GGML_TYPE_F16:
4171        case GGML_TYPE_COUNT:
4172            {
4173                assert(false);
4174            } break;
4175    }
4176}
4177
4178// ggml_compute_forward_mul_mat
4179
4180#if defined(GGML_USE_ACCELERATE) || defined(GGML_USE_OPENBLAS)
4181// helper function to determine if it is better to use BLAS or not
4182// for large matrices, BLAS is faster
4183static bool ggml_compute_forward_mul_mat_use_blas(
4184        const struct ggml_tensor * src0,
4185        const struct ggml_tensor * src1,
4186              struct ggml_tensor * dst) {
4187    UNUSED(src0);
4188
4189    const int ne10 = src1->ne[0];
4190
4191    const int ne0 = dst->ne[0];
4192    const int ne1 = dst->ne[1];
4193
4194    // TODO: find the optimal values for these
4195    if (ggml_is_contiguous(src0) && ggml_is_contiguous(src1) && ne0 >= 32 && ne1 >= 32 && ne10 >= 32) {
4196        //printf("BLAS: %d %d %d\n", ne0, ne1, ne10);
4197        return true;
4198    }
4199
4200    return false;
4201}
4202#endif
4203
4204static void ggml_compute_forward_mul_mat_f32(
4205        const struct ggml_compute_params * params,
4206        const struct ggml_tensor * src0,
4207        const struct ggml_tensor * src1,
4208              struct ggml_tensor * dst) {
4209    int64_t t0 = ggml_perf_time_us();
4210    UNUSED(t0);
4211
4212    const int ne00 = src0->ne[0];
4213    const int ne01 = src0->ne[1];
4214    const int ne02 = src0->ne[2];
4215    const int ne03 = src0->ne[3];
4216
4217    const int ne10 = src1->ne[0];
4218    const int ne11 = src1->ne[1];
4219    const int ne12 = src1->ne[2];
4220    const int ne13 = src1->ne[3];
4221
4222    const int ne0  = dst->ne[0];
4223    const int ne1  = dst->ne[1];
4224    const int ne2  = dst->ne[2];
4225    const int ne3  = dst->ne[3];
4226    const int ne   = ne0*ne1*ne2*ne3;
4227
4228    const int nb00 = src0->nb[0];
4229    const int nb01 = src0->nb[1];
4230    const int nb02 = src0->nb[2];
4231    const int nb03 = src0->nb[3];
4232
4233    const int nb10 = src1->nb[0];
4234    const int nb11 = src1->nb[1];
4235    const int nb12 = src1->nb[2];
4236    const int nb13 = src1->nb[3];
4237
4238    const int nb0  = dst->nb[0];
4239    const int nb1  = dst->nb[1];
4240    const int nb2  = dst->nb[2];
4241    const int nb3  = dst->nb[3];
4242
4243    const int ith = params->ith;
4244    const int nth = params->nth;
4245
4246    assert(ne02 == ne12);
4247    assert(ne03 == ne13);
4248    assert(ne2  == ne12);
4249    assert(ne3  == ne13);
4250
4251    // TODO: we don't support permuted src0
4252    assert(nb00 == sizeof(float) || nb01 == sizeof(float));
4253
4254    // dst cannot be transposed or permuted
4255    assert(nb0 == sizeof(float));
4256    assert(nb0 <= nb1);
4257    assert(nb1 <= nb2);
4258    assert(nb2 <= nb3);
4259
4260    assert(ne0 == ne01);
4261    assert(ne1 == ne11);
4262    assert(ne2 == ne02);
4263    assert(ne3 == ne03);
4264
4265    // nb01 >= nb00 - src0 is not transposed
4266    //   compute by src0 rows
4267    //
4268    // nb00 <  nb01 - src0 is transposed
4269    //   compute by src0 columns
4270
4271#if defined(GGML_USE_ACCELERATE) || defined(GGML_USE_OPENBLAS)
4272    if (ggml_compute_forward_mul_mat_use_blas(src0, src1, dst)) {
4273        GGML_ASSERT(nb10 == sizeof(float));
4274
4275        if (params->ith != 0) return;
4276
4277        if (params->type == GGML_TASK_INIT) {
4278            return;
4279        }
4280
4281        if (params->type == GGML_TASK_FINALIZE) {
4282            return;
4283        }
4284
4285        for (int i03 = 0; i03 < ne03; i03++) {
4286            for (int i02 = 0; i02 < ne02; i02++) {
4287                const float * x = (float *) (src0->data);
4288                const float * y = (float *) ((char *) src1->data + i02*nb12 + i03*nb13);
4289
4290                float * d = (float *) ((char *) dst->data + i02*nb2 + i03*nb3);
4291
4292                // zT = y * xT
4293                {
4294                    cblas_sgemm(CblasRowMajor, CblasNoTrans, CblasTrans,
4295                            ne11, ne01, ne10,
4296                            1.0f,    y, ne10,
4297                                     x, ne10,
4298                            0.0f,    d, ne01);
4299                }
4300            }
4301        }
4302
4303        //printf("CBLAS F32 = %f ms, %d x %d x %d x %d\n", (ggml_perf_time_us() - t0)/1000.0, ne0, ne1, ne2, ne3);
4304
4305        return;
4306    }
4307#endif
4308
4309    if (params->type == GGML_TASK_INIT) {
4310        if (nb01 >= nb00) {
4311            return;
4312        }
4313
4314        // TODO: fix this memset (wsize is overestimated)
4315        memset(params->wdata, 0, params->wsize);
4316        return;
4317    }
4318
4319    if (params->type == GGML_TASK_FINALIZE) {
4320        if (nb01 >= nb00) {
4321            return;
4322        }
4323
4324        // TODO: fix this memset (wsize is overestimated)
4325        //assert(params->wsize == (ggml_nbytes(dst) + CACHE_LINE_SIZE)*nth);
4326
4327        float * const wdata = params->wdata;
4328
4329        // cols per thread
4330        const int dc = (ne + nth - 1)/nth;
4331
4332        // col range for this thread
4333        const int ic0 = dc*ith;
4334        const int ic1 = MIN(ic0 + dc, ne);
4335
4336        ggml_vec_cpy_f32(ic1 - ic0, (float *) dst->data + ic0, wdata + ic0);
4337
4338        for (int k = 1; k < nth; k++) {
4339            ggml_vec_acc_f32(ic1 - ic0, (float *) dst->data + ic0, wdata + (ne + CACHE_LINE_SIZE_F32)*k + ic0);
4340        }
4341
4342        return;
4343    }
4344
4345    if (nb01 >= nb00) {
4346        // TODO: do not support transposed src1
4347        assert(nb10 == sizeof(float));
4348
4349        // parallelize by src0 rows using ggml_vec_dot_f32
4350
4351        // total rows in src0
4352        const int nr = ne01*ne02*ne03;
4353
4354        // rows per thread
4355        const int dr = (nr + nth - 1)/nth;
4356
4357        // row range for this thread
4358        const int ir0 = dr*ith;
4359        const int ir1 = MIN(ir0 + dr, nr);
4360
4361        for (int ir = ir0; ir < ir1; ++ir) {
4362            // src0 indices
4363            const int i03 = ir/(ne02*ne01);
4364            const int i02 = (ir - i03*ne02*ne01)/ne01;
4365            const int i01 = (ir - i03*ne02*ne01 - i02*ne01);
4366
4367            for (int ic = 0; ic < ne11; ++ic) {
4368                // src1 indices
4369                const int i13 = i03;
4370                const int i12 = i02;
4371                const int i11 = ic;
4372
4373                // dst indices
4374                const int i0 = i01;
4375                const int i1 = i11;
4376                const int i2 = i02;
4377                const int i3 = i03;
4378
4379                ggml_vec_dot_f32(ne00,
4380                        (float *) ((char *)  dst->data + (i0*nb0 + i1*nb1 + i2*nb2 + i3*nb3)),
4381                        (float *) ((char *) src0->data + (i01*nb01 + i02*nb02 + i03*nb03)),
4382                        (float *) ((char *) src1->data + (i11*nb11 + i12*nb12 + i13*nb13)));
4383            }
4384        }
4385    } else {
4386        // parallelize by src1 columns using ggml_vec_mad_f32
4387        // each thread has its own work data
4388        // during FINALIZE we accumulate all work data into dst
4389
4390        // total columns in src1
4391        const int nc = ne10;
4392
4393        // columns per thread
4394        const int dc = (nc + nth - 1)/nth;
4395
4396        // column range for this thread
4397        const int ic0 = dc*ith;
4398        const int ic1 = MIN(ic0 + dc, nc);
4399
4400        // work data for thread
4401        const int wo = (ne + CACHE_LINE_SIZE_F32)*ith;
4402        float * const wdata = params->wdata;
4403
4404        for (int i13 = 0; i13 < ne13; ++i13) {
4405            for (int i12 = 0; i12 < ne12; ++i12) {
4406                for (int i11 = 0; i11 < ne11; ++i11) {
4407                    for (int ic = ic0; ic < ic1; ++ic) {
4408                        // src1 indices
4409                        const int i10 = ic;
4410
4411                        // src0 indices
4412                        const int i03 = i13;
4413                        const int i02 = i12;
4414                        const int i00 = ic;
4415
4416                        // dst indices
4417                        const int i1 = i11;
4418                        const int i2 = i12;
4419                        const int i3 = i13;
4420
4421                        assert(sizeof(float)*(wo + i3*ne2*ne1*ne0 + i2*ne1*ne0 + i1*ne0 + ne01) <= params->wsize);
4422
4423                        ggml_vec_mad_f32(ne01,
4424                                (float *) (wdata + wo + i3*ne2*ne1*ne0 + i2*ne1*ne0 + i1*ne0),
4425                                (float *) ((char *) src0->data + (i00*nb00 + i02*nb02 + i03*nb03)),
4426                               *(float *) ((char *) src1->data + (i10*nb10 + i11*nb11 + i12*nb12 + i13*nb13)));
4427                    }
4428                }
4429            }
4430        }
4431    }
4432
4433    //int64_t t1 = ggml_perf_time_us();
4434    //static int64_t acc = 0;
4435    //acc += t1 - t0;
4436    //if (t1 - t0 > 10) {
4437    //    printf("\n");
4438    //    printf("ne00 = %5d, ne01 = %5d, ne02 = %5d, ne03 = %5d\n", ne00, ne01, ne02, ne03);
4439    //    printf("nb00 = %5d, nb01 = %5d, nb02 = %5d, nb03 = %5d\n", nb00, nb01, nb02, nb03);
4440    //    printf("ne10 = %5d, ne11 = %5d, ne12 = %5d, ne13 = %5d\n", ne10, ne11, ne12, ne13);
4441    //    printf("nb10 = %5d, nb11 = %5d, nb12 = %5d, nb13 = %5d\n", nb10, nb11, nb12, nb13);
4442
4443    //    printf("XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX task %d/%d: %d us, acc = %d\n", ith, nth, (int) (t1 - t0), (int) acc);
4444    //}
4445}
4446
4447static void ggml_compute_forward_mul_mat_f16_f32(
4448        const struct ggml_compute_params * params,
4449        const struct ggml_tensor * src0,
4450        const struct ggml_tensor * src1,
4451              struct ggml_tensor * dst) {
4452    int64_t t0 = ggml_perf_time_us();
4453    UNUSED(t0);
4454
4455    const int ne00 = src0->ne[0];
4456    const int ne01 = src0->ne[1];
4457    const int ne02 = src0->ne[2];
4458    const int ne03 = src0->ne[3];
4459
4460    const int ne10 = src1->ne[0];
4461    const int ne11 = src1->ne[1];
4462    const int ne12 = src1->ne[2];
4463    const int ne13 = src1->ne[3];
4464
4465    const int ne0  = dst->ne[0];
4466    const int ne1  = dst->ne[1];
4467    const int ne2  = dst->ne[2];
4468    const int ne3  = dst->ne[3];
4469    const int ne   = ne0*ne1*ne2*ne3;
4470
4471    const int nb00 = src0->nb[0];
4472    const int nb01 = src0->nb[1];
4473    const int nb02 = src0->nb[2];
4474    const int nb03 = src0->nb[3];
4475
4476    const int nb10 = src1->nb[0];
4477    const int nb11 = src1->nb[1];
4478    const int nb12 = src1->nb[2];
4479    const int nb13 = src1->nb[3];
4480
4481    const int nb0  = dst->nb[0];
4482    const int nb1  = dst->nb[1];
4483    const int nb2  = dst->nb[2];
4484    const int nb3  = dst->nb[3];
4485
4486    const int ith = params->ith;
4487    const int nth = params->nth;
4488
4489    GGML_ASSERT(ne02 == ne12);
4490    GGML_ASSERT(ne03 == ne13);
4491    GGML_ASSERT(ne2  == ne12);
4492    GGML_ASSERT(ne3  == ne13);
4493
4494    // TODO: we don't support permuted src0
4495    GGML_ASSERT(nb00 == sizeof(ggml_fp16_t) || nb01 == sizeof(ggml_fp16_t));
4496
4497    // dst cannot be transposed or permuted
4498    GGML_ASSERT(nb0 == sizeof(float));
4499    GGML_ASSERT(nb0 <= nb1);
4500    GGML_ASSERT(nb1 <= nb2);
4501    GGML_ASSERT(nb2 <= nb3);
4502
4503    GGML_ASSERT(ne0 == ne01);
4504    GGML_ASSERT(ne1 == ne11);
4505    GGML_ASSERT(ne2 == ne02);
4506    GGML_ASSERT(ne3 == ne03);
4507
4508    // nb01 >= nb00 - src0 is not transposed
4509    //   compute by src0 rows
4510    //
4511    // nb00 <  nb01 - src0 is transposed
4512    //   compute by src0 columns
4513
4514#if defined(GGML_USE_ACCELERATE) || defined(GGML_USE_OPENBLAS)
4515    if (ggml_compute_forward_mul_mat_use_blas(src0, src1, dst)) {
4516        GGML_ASSERT(nb10 == sizeof(float));
4517
4518        if (params->ith != 0) return;
4519
4520        if (params->type == GGML_TASK_INIT) {
4521            return;
4522        }
4523
4524        if (params->type == GGML_TASK_FINALIZE) {
4525            return;
4526        }
4527
4528        float * const wdata = params->wdata;
4529
4530        for (int i03 = 0; i03 < ne03; i03++) {
4531            for (int i02 = 0; i02 < ne02; i02++) {
4532                {
4533                    int id = 0;
4534                    for (int i01 = 0; i01 < ne01; ++i01) {
4535                        for (int i00 = 0; i00 < ne00; ++i00) {
4536                            wdata[id++] = GGML_FP16_TO_FP32(*(ggml_fp16_t *) ((char *) src0->data + i03*nb03 + i02*nb02 + i01*nb01 + i00*nb00));
4537                        }
4538                    }
4539                }
4540
4541                const float * x = wdata;
4542                const float * y = (float *) ((char *) src1->data + i02*nb12 + i03*nb13);
4543
4544                //      float * z =                          wdata + ne00*ne01;
4545
4546                // z = x * yT
4547                //{
4548                //    cblas_sgemm(CblasRowMajor, CblasNoTrans, CblasTrans,
4549                //            ne01, ne11, ne00,
4550                //            1.0f, x, ne00,
4551                //                  y, ne00,
4552                //            0.0f, z, ne11);
4553                //}
4554
4555                float * d = (float *) ((char *) dst->data + i02*nb2 + i03*nb3);
4556
4557                // transpose z
4558                //for (int j = 0; j < ne11; ++j) {
4559                //    for (int i = 0; i < ne01; ++i) {
4560                //        d[j*ne01 + i] = z[i*ne11 + j];
4561                //    }
4562                //}
4563
4564                {
4565#if 1
4566                    // zT = y * xT
4567                    cblas_sgemm(CblasRowMajor, CblasNoTrans, CblasTrans,
4568                            ne11, ne01, ne10,
4569                            1.0f,    y, ne00,
4570                                     x, ne00,
4571                            0.0f,    d, ne01);
4572#else
4573                    // zT = (xT * y)T
4574                    cblas_sgemm(CblasColMajor, CblasTrans, CblasNoTrans,
4575                            ne01, ne11, ne10,
4576                            1.0f,    x, ne00,
4577                                     y, ne00,
4578                            0.0f,    d, ne01);
4579#endif
4580                }
4581            }
4582        }
4583
4584        //printf("CBLAS = %f ms, %d x %d x %d x %d\n", (ggml_perf_time_us() - t0)/1000.0, ne0, ne1, ne2, ne3);
4585
4586        return;
4587    }
4588#endif
4589
4590    if (params->type == GGML_TASK_INIT) {
4591        if (nb01 >= nb00) {
4592            ggml_fp16_t * const wdata = params->wdata;
4593
4594            int id = 0;
4595            for (int i13 = 0; i13 < ne13; ++i13) {
4596                for (int i12 = 0; i12 < ne12; ++i12) {
4597                    for (int i11 = 0; i11 < ne11; ++i11) {
4598                        for (int i10 = 0; i10 < ne10; ++i10) {
4599                            wdata[id++] = GGML_FP32_TO_FP16(*(float *)((char *) src1->data + i13*nb13 + i12*nb12 + i11*nb11 + i10*nb10));
4600                        }
4601                    }
4602                }
4603            }
4604
4605            GGML_ASSERT(id*sizeof(ggml_fp16_t) <= params->wsize);
4606
4607            return;
4608        }
4609
4610        // TODO: fix this memset (wsize is overestimated)
4611        memset(params->wdata, 0, params->wsize);
4612        return;
4613    }
4614
4615    if (params->type == GGML_TASK_FINALIZE) {
4616        if (nb01 >= nb00) {
4617            return;
4618        }
4619
4620        // TODO: fix this memset (wsize is overestimated)
4621        //assert(params->wsize == (ggml_nbytes(dst) + CACHE_LINE_SIZE)*nth);
4622
4623        ggml_fp16_t * const wdata = params->wdata;
4624
4625        // cols per thread
4626        const int dc = (ne + nth - 1)/nth;
4627
4628        // col range for this thread
4629        const int ic0 = dc*ith;
4630        const int ic1 = MIN(ic0 + dc, ne);
4631
4632        for (int i = ic0; i < ic1; ++i) {
4633            ((float *) dst->data)[i] = GGML_FP16_TO_FP32(wdata[i]);
4634        }
4635
4636        for (int k = 1; k < nth; k++) {
4637            for (int i = ic0; i < ic1; ++i) {
4638                ((float *) dst->data)[i] += GGML_FP16_TO_FP32(wdata[(ne + CACHE_LINE_SIZE_F32)*k + i]);
4639            }
4640        }
4641
4642        return;
4643    }
4644
4645    if (nb01 >= nb00) {
4646        // fp16 -> half the size, so divide by 2
4647        // TODO: do not support transposed src1
4648        assert(nb10/2 == sizeof(ggml_fp16_t));
4649
4650        // parallelize by src0 rows using ggml_vec_dot_f32
4651
4652        // total rows in src0
4653        const int nr = ne01*ne02*ne03;
4654
4655        // rows per thread
4656        const int dr = (nr + nth - 1)/nth;
4657
4658        // row range for this thread
4659        const int ir0 = dr*ith;
4660        const int ir1 = MIN(ir0 + dr, nr);
4661
4662        ggml_fp16_t * wdata = params->wdata;
4663
4664        for (int ir = ir0; ir < ir1; ++ir) {
4665            // src0 indices
4666            const int i03 = ir/(ne02*ne01);
4667            const int i02 = (ir - i03*ne02*ne01)/ne01;
4668            const int i01 = (ir - i03*ne02*ne01 - i02*ne01);
4669
4670            const int i13 = i03;
4671            const int i12 = i02;
4672
4673            const int i0 = i01;
4674            const int i2 = i02;
4675            const int i3 = i03;
4676
4677            ggml_fp16_t * src0_row = (ggml_fp16_t *) ((char *) src0->data + (i01*nb01 + i02*nb02 + i03*nb03));
4678            ggml_fp16_t * src1_col = wdata + (i13*ne12*ne11 + i12*ne11 + 0)*ne00;
4679
4680            float * dst_col = (float *) ((char *) dst->data + (i0*nb0 + 0*nb1 + i2*nb2 + i3*nb3));
4681
4682            for (int ic = 0; ic < ne11; ++ic) {
4683                assert(ne00 % 32 == 0);
4684
4685                ggml_vec_dot_f16(ne00, &dst_col[ic*ne0], src0_row, src1_col + ic*ne00);
4686            }
4687        }
4688    } else {
4689        // parallelize by src1 columns using ggml_vec_mad_f32
4690        // each thread has its own work data
4691        // during FINALIZE we accumulate all work data into dst
4692
4693        // total columns in src1
4694        const int nc = ne10;
4695
4696        // columns per thread
4697        const int dc = (nc + nth - 1)/nth;
4698
4699        // column range for this thread
4700        const int ic0 = dc*ith;
4701        const int ic1 = MIN(ic0 + dc, nc);
4702
4703        // work data for thread
4704        const int wo = (ne + CACHE_LINE_SIZE_F32)*ith;
4705        ggml_fp16_t * const wdata = params->wdata;
4706
4707        for (int i13 = 0; i13 < ne13; ++i13) {
4708            for (int i12 = 0; i12 < ne12; ++i12) {
4709                for (int i11 = 0; i11 < ne11; ++i11) {
4710                    // dst indices
4711                    const int i1 = i11;
4712                    const int i2 = i12;
4713                    const int i3 = i13;
4714
4715                    ggml_fp16_t * dst_row = wdata + wo + i3*ne2*ne1*ne0 + i2*ne1*ne0 + i1*ne0;
4716
4717                    for (int ic = ic0; ic < ic1; ++ic) {
4718                        // src1 indices
4719                        const int i10 = ic;
4720
4721                        // src0 indices
4722                        const int i03 = i13;
4723                        const int i02 = i12;
4724                        const int i00 = ic;
4725
4726                        assert(sizeof(ggml_fp16_t)*(wo + i3*ne2*ne1*ne0 + i2*ne1*ne0 + i1*ne0 + ne01) <= params->wsize);
4727
4728                        ggml_fp16_t * src0_col =  (ggml_fp16_t *) ((char *) src0->data + (i00*nb00 + i02*nb02 + i03*nb03));
4729                        float         src1_val = *      (float *) ((char *) src1->data + (i10*nb10 + i11*nb11 + i12*nb12 + i13*nb13));
4730
4731                        ggml_vec_mad_f16(ne01, dst_row, src0_col, src1_val);
4732                    }
4733                }
4734            }
4735        }
4736    }
4737
4738    //int64_t t1 = ggml_time_us();
4739    //static int64_t acc = 0;
4740    //acc += t1 - t0;
4741    //if (t1 - t0 > 10) {
4742    //    printf("\n");
4743    //    printf("ne00 = %5d, ne01 = %5d, ne02 = %5d, ne03 = %5d\n", ne00, ne01, ne02, ne03);
4744    //    printf("nb00 = %5d, nb01 = %5d, nb02 = %5d, nb03 = %5d\n", nb00, nb01, nb02, nb03);
4745    //    printf("ne10 = %5d, ne11 = %5d, ne12 = %5d, ne13 = %5d\n", ne10, ne11, ne12, ne13);
4746
4747    //    printf("XXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXXX task %d/%d: %d us, acc = %d\n", ith, nth, (int) (t1 - t0), (int) acc);
4748    //}
4749}
4750
4751static void ggml_compute_forward_mul_mat(
4752        const struct ggml_compute_params * params,
4753        const struct ggml_tensor * src0,
4754        const struct ggml_tensor * src1,
4755        struct ggml_tensor * dst) {
4756    switch (src0->type) {
4757        case GGML_TYPE_F16:
4758            {
4759                ggml_compute_forward_mul_mat_f16_f32(params, src0, src1, dst);
4760            } break;
4761        case GGML_TYPE_F32:
4762            {
4763                ggml_compute_forward_mul_mat_f32(params, src0, src1, dst);
4764            } break;
4765        case GGML_TYPE_I8:
4766        case GGML_TYPE_I16:
4767        case GGML_TYPE_I32:
4768        case GGML_TYPE_COUNT:
4769            {
4770                assert(false);
4771            } break;
4772    }
4773}
4774
4775// ggml_compute_forward_scale
4776
4777static void ggml_compute_forward_scale_f32(
4778        const struct ggml_compute_params * params,
4779        const struct ggml_tensor * src0,
4780        const struct ggml_tensor * src1,
4781        struct ggml_tensor * dst) {
4782    GGML_ASSERT(ggml_is_contiguous(src0));
4783    GGML_ASSERT(ggml_is_contiguous(dst));
4784    GGML_ASSERT(ggml_are_same_shape(src0, dst));
4785    GGML_ASSERT(ggml_is_scalar(src1));
4786
4787    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
4788        return;
4789    }
4790
4791    // scale factor
4792    const float v = *(float *) src1->data;
4793
4794    const int ith = params->ith;
4795    const int nth = params->nth;
4796
4797    const int nc = src0->ne[0];
4798    const int nr = ggml_nrows(src0);
4799
4800    // rows per thread
4801    const int dr = (nr + nth - 1)/nth;
4802
4803    // row range for this thread
4804    const int ir0 = dr*ith;
4805    const int ir1 = MIN(ir0 + dr, nr);
4806
4807    for (int i1 = ir0; i1 < ir1; i1++) {
4808        ggml_vec_scale_f32(nc, (float *) ((char *) dst->data + i1*(dst->nb[1])), v);
4809    }
4810}
4811
4812static void ggml_compute_forward_scale(
4813        const struct ggml_compute_params * params,
4814        const struct ggml_tensor * src0,
4815        const struct ggml_tensor * src1,
4816        struct ggml_tensor * dst) {
4817    switch (src0->type) {
4818        case GGML_TYPE_F32:
4819            {
4820                ggml_compute_forward_scale_f32(params, src0, src1, dst);
4821            } break;
4822        case GGML_TYPE_I8:
4823        case GGML_TYPE_I16:
4824        case GGML_TYPE_I32:
4825        case GGML_TYPE_F16:
4826        case GGML_TYPE_COUNT:
4827            {
4828                assert(false);
4829            } break;
4830    }
4831}
4832
4833// ggml_compute_forward_cpy
4834
4835static void ggml_compute_forward_cpy(
4836        const struct ggml_compute_params * params,
4837        const struct ggml_tensor * src0,
4838        struct ggml_tensor * dst) {
4839    ggml_compute_forward_dup(params, src0, dst);
4840}
4841
4842// ggml_compute_forward_reshape
4843
4844static void ggml_compute_forward_reshape(
4845        const struct ggml_compute_params * params,
4846        const struct ggml_tensor * src0,
4847        struct ggml_tensor * dst) {
4848    // NOP
4849    UNUSED(params);
4850    UNUSED(src0);
4851    UNUSED(dst);
4852}
4853
4854// ggml_compute_forward_view
4855
4856static void ggml_compute_forward_view(
4857        const struct ggml_compute_params * params,
4858        const struct ggml_tensor * src0) {
4859    // NOP
4860    UNUSED(params);
4861    UNUSED(src0);
4862}
4863
4864// ggml_compute_forward_permute
4865
4866static void ggml_compute_forward_permute(
4867        const struct ggml_compute_params * params,
4868        const struct ggml_tensor * src0) {
4869    // NOP
4870    UNUSED(params);
4871    UNUSED(src0);
4872}
4873
4874// ggml_compute_forward_transpose
4875
4876static void ggml_compute_forward_transpose(
4877        const struct ggml_compute_params * params,
4878        const struct ggml_tensor * src0) {
4879    // NOP
4880    UNUSED(params);
4881    UNUSED(src0);
4882}
4883
4884// ggml_compute_forward_get_rows
4885
4886static void ggml_compute_forward_get_rows_f16(
4887        const struct ggml_compute_params * params,
4888        const struct ggml_tensor * src0,
4889        const struct ggml_tensor * src1,
4890              struct ggml_tensor * dst) {
4891    assert(params->ith == 0);
4892
4893    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
4894        return;
4895    }
4896
4897    const int nc = src0->ne[0];
4898    const int nr = ggml_nelements(src1);
4899
4900    assert( dst->ne[0] == nc);
4901    assert( dst->ne[1] == nr);
4902    assert(src0->nb[0] == sizeof(ggml_fp16_t));
4903
4904    for (int i = 0; i < nr; ++i) {
4905        const int r = ((int32_t *) src1->data)[i];
4906
4907        for (int j = 0; j < nc; ++j) {
4908            ggml_fp16_t v = ((ggml_fp16_t *) ((char *) src0->data + r*src0->nb[1]))[j];
4909            ((float *) ((char *)  dst->data + i*dst->nb[1]))[j] = GGML_FP16_TO_FP32(v);
4910        }
4911    }
4912}
4913
4914static void ggml_compute_forward_get_rows_f32(
4915        const struct ggml_compute_params * params,
4916        const struct ggml_tensor * src0,
4917        const struct ggml_tensor * src1,
4918              struct ggml_tensor * dst) {
4919    assert(params->ith == 0);
4920
4921    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
4922        return;
4923    }
4924
4925    const int nc = src0->ne[0];
4926    const int nr = ggml_nelements(src1);
4927
4928    assert( dst->ne[0] == nc);
4929    assert( dst->ne[1] == nr);
4930    assert(src0->nb[0] == sizeof(float));
4931
4932    for (int i = 0; i < nr; ++i) {
4933        const int r = ((int32_t *) src1->data)[i];
4934
4935        ggml_vec_cpy_f32(nc,
4936                (float *) ((char *)  dst->data + i*dst->nb[1]),
4937                (float *) ((char *) src0->data + r*src0->nb[1]));
4938    }
4939}
4940
4941static void ggml_compute_forward_get_rows(
4942        const struct ggml_compute_params * params,
4943        const struct ggml_tensor * src0,
4944        const struct ggml_tensor * src1,
4945        struct ggml_tensor * dst) {
4946    switch (src0->type) {
4947        case GGML_TYPE_F16:
4948            {
4949                ggml_compute_forward_get_rows_f16(params, src0, src1, dst);
4950            } break;
4951        case GGML_TYPE_F32:
4952            {
4953                ggml_compute_forward_get_rows_f32(params, src0, src1, dst);
4954            } break;
4955        case GGML_TYPE_I8:
4956        case GGML_TYPE_I16:
4957        case GGML_TYPE_I32:
4958        case GGML_TYPE_COUNT:
4959            {
4960                assert(false);
4961            } break;
4962    }
4963}
4964
4965// ggml_compute_forward_diag_mask_inf
4966
4967static void ggml_compute_forward_diag_mask_inf_f32(
4968        const struct ggml_compute_params * params,
4969        const struct ggml_tensor * src0,
4970        const struct ggml_tensor * src1,
4971        struct ggml_tensor * dst) {
4972    assert(params->ith == 0);
4973    assert(src1->type == GGML_TYPE_I32);
4974    assert(ggml_nelements(src1) == 1);
4975
4976    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
4977        return;
4978    }
4979
4980    const int n_past = ((int32_t *) src1->data)[0];
4981
4982    // TODO: handle transposed/permuted matrices
4983
4984    const int n  = ggml_nrows(src0);
4985    const int nc = src0->ne[0];
4986    const int nr = src0->ne[1];
4987    const int nz = n/nr;
4988
4989    assert( dst->nb[0] == sizeof(float));
4990    assert(src0->nb[0] == sizeof(float));
4991
4992    for (int k = 0; k < nz; k++) {
4993        for (int j = 0; j < nr; j++) {
4994            for (int i = n_past; i < nc; i++) {
4995                if (i > n_past + j) {
4996                    *(float *)((char *) dst->data + k*dst->nb[2] + j*dst->nb[1] + i*dst->nb[0]) = -INFINITY;
4997                }
4998            }
4999        }
5000    }
5001}
5002
5003static void ggml_compute_forward_diag_mask_inf(
5004        const struct ggml_compute_params * params,
5005        const struct ggml_tensor * src0,
5006        const struct ggml_tensor * src1,
5007        struct ggml_tensor * dst) {
5008    switch (src0->type) {
5009        case GGML_TYPE_F32:
5010            {
5011                ggml_compute_forward_diag_mask_inf_f32(params, src0, src1, dst);
5012            } break;
5013        case GGML_TYPE_I8:
5014        case GGML_TYPE_I16:
5015        case GGML_TYPE_I32:
5016        case GGML_TYPE_F16:
5017        case GGML_TYPE_COUNT:
5018            {
5019                assert(false);
5020            } break;
5021    }
5022}
5023
5024// ggml_compute_forward_soft_max
5025
5026static void ggml_compute_forward_soft_max_f32(
5027        const struct ggml_compute_params * params,
5028        const struct ggml_tensor * src0,
5029        struct ggml_tensor * dst) {
5030    GGML_ASSERT(ggml_is_contiguous(src0));
5031    GGML_ASSERT(ggml_is_contiguous(dst));
5032    GGML_ASSERT(ggml_are_same_shape(src0, dst));
5033
5034    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
5035        return;
5036    }
5037
5038    // TODO: handle transposed/permuted matrices
5039
5040    const int ith = params->ith;
5041    const int nth = params->nth;
5042
5043    const int nc = src0->ne[0];
5044    const int nr = ggml_nrows(src0);
5045
5046    // rows per thread
5047    const int dr = (nr + nth - 1)/nth;
5048
5049    // row range for this thread
5050    const int ir0 = dr*ith;
5051    const int ir1 = MIN(ir0 + dr, nr);
5052
5053    for (int i1 = ir0; i1 < ir1; i1++) {
5054        float *p = (float *)((char *) dst->data + i1*dst->nb[1]);
5055
5056#ifndef NDEBUG
5057        for (int i = 0; i < nc; ++i) {
5058            assert(!isnan(p[i]));
5059        }
5060#endif
5061
5062        float max = -INFINITY;
5063        for (int i = 0; i < nc; i++) {
5064            max = MAX(max, p[i]);
5065        }
5066
5067        ggml_float sum = 0.0;
5068
5069        uint16_t ss;
5070        for (int i = 0; i < nc; i++) {
5071            if (p[i] == -INFINITY) {
5072                p[i] = 0.0;
5073            } else {
5074                //const float val = (p[i] == -INFINITY) ? 0.0 : exp(p[i] - max);
5075                ggml_fp16_t s = GGML_FP32_TO_FP16(p[i] - max);
5076                memcpy(&ss, &s, sizeof(ss));
5077                const float val = GGML_FP16_TO_FP32(table_exp_f16[ss]);
5078                sum += val;
5079                p[i] = val;
5080            }
5081        }
5082
5083        assert(sum > 0.0f);
5084
5085        sum = 1.0/sum;
5086        ggml_vec_scale_f32(nc, p, sum);
5087
5088#ifndef NDEBUG
5089        for (int i = 0; i < nc; ++i) {
5090            assert(!isnan(p[i]));
5091            assert(!isinf(p[i]));
5092        }
5093#endif
5094    }
5095}
5096
5097static void ggml_compute_forward_soft_max(
5098        const struct ggml_compute_params * params,
5099        const struct ggml_tensor * src0,
5100        struct ggml_tensor * dst) {
5101    switch (src0->type) {
5102        case GGML_TYPE_F32:
5103            {
5104                ggml_compute_forward_soft_max_f32(params, src0, dst);
5105            } break;
5106        case GGML_TYPE_I8:
5107        case GGML_TYPE_I16:
5108        case GGML_TYPE_I32:
5109        case GGML_TYPE_F16:
5110        case GGML_TYPE_COUNT:
5111            {
5112                assert(false);
5113            } break;
5114    }
5115}
5116
5117// ggml_compute_forward_rope
5118
5119static void ggml_compute_forward_rope_f32(
5120        const struct ggml_compute_params * params,
5121        const struct ggml_tensor * src0,
5122        const struct ggml_tensor * src1,
5123        struct ggml_tensor * dst) {
5124    assert(params->ith == 0);
5125    assert(src1->type == GGML_TYPE_I32);
5126    assert(ggml_nelements(src1) == 3);
5127
5128    if (params->type == GGML_TASK_INIT || params->type == GGML_TASK_FINALIZE) {
5129        return;
5130    }
5131
5132    const int n_past = ((int32_t *) src1->data)[0];
5133    const int n_dims = ((int32_t *) src1->data)[1];
5134    const int mode   = ((int32_t *) src1->data)[2];
5135
5136    //const int ne0 = src0->ne[0];
5137    const int ne1 = src0->ne[1];
5138    const int ne2 = src0->ne[2];
5139    const int ne3 = src0->ne[3];
5140
5141    const int nb0 = src0->nb[0];
5142    const int nb1 = src0->nb[1];
5143    const int nb2 = src0->nb[2];
5144    const int nb3 = src0->nb[3];
5145
5146    //printf("ne0: %d, ne1: %d, ne2: %d, ne3: %d\n", ne0, ne1, ne2, ne3);
5147    //printf("n_past = %d, ne2 = %d\n", n_past, ne2);
5148
5149    assert(nb0 == sizeof(float));
5150
5151    // TODO: optimize
5152    for (int i3 = 0; i3 < ne3; i3++) {
5153        for (int i2 = (mode == 0 ? 0 : n_past); i2 < ne2; i2++) {
5154            const int p = (mode == 0 ? n_past + i2 : i2);
5155            for (int i1 = 0; i1 < ne1; i1++) {
5156                for (int i0 = 0; i0 < n_dims; i0 += 2) {
5157                    const double theta = pow(10000.0, ((double)-i0)/n_dims);
5158
5159                    const double cos_theta = cos(p*theta);
5160                    const double sin_theta = sin(p*theta);
5161
5162                    const float * const src = (float *)((char *) src0->data + i3*nb3 + i2*nb2 + i1*nb1 + i0*nb0);
5163                          float * dst_data  = (float *)((char *)  dst->data + i3*nb3 + i2*nb2 + i1*nb1 + i0*nb0);
5164
5165                    double x0 = src[0];
5166                    double x1 = src[1];
5167
5168                    dst_data[0] = x0*cos_theta - x1*sin_theta;
5169                    dst_data[1] = x0*sin_theta + x1*cos_theta;
5170                }
5171            }
5172        }
5173    }
5174}
5175
5176static void ggml_compute_forward_rope(
5177        const struct ggml_compute_params * params,
5178        const struct ggml_tensor * src0,
5179        const struct ggml_tensor * src1,
5180        struct ggml_tensor * dst) {
5181    switch (src0->type) {
5182        case GGML_TYPE_F32:
5183            {
5184                ggml_compute_forward_rope_f32(params, src0, src1, dst);
5185            } break;
5186        case GGML_TYPE_I8:
5187        case GGML_TYPE_I16:
5188        case GGML_TYPE_I32:
5189        case GGML_TYPE_F16:
5190        case GGML_TYPE_COUNT:
5191            {
5192                assert(false);
5193            } break;
5194    }
5195}
5196
5197// ggml_compute_forward_conv_1d_1s
5198
5199static void ggml_compute_forward_conv_1d_1s_f16_f32(
5200        const struct ggml_compute_params * params,
5201        const struct ggml_tensor * src0,
5202        const struct ggml_tensor * src1,
5203              struct ggml_tensor * dst) {
5204    GGML_ASSERT(src0->type == GGML_TYPE_F16);
5205    GGML_ASSERT(src1->type == GGML_TYPE_F32);
5206    GGML_ASSERT( dst->type == GGML_TYPE_F32);
5207
5208    int64_t t0 = ggml_perf_time_us();
5209    UNUSED(t0);
5210
5211    const int ne00 = src0->ne[0];
5212    const int ne01 = src0->ne[1];
5213    const int ne02 = src0->ne[2];
5214    //const int ne03 = src0->ne[3];
5215
5216    const int ne10 = src1->ne[0];
5217    const int ne11 = src1->ne[1];
5218    //const int ne12 = src1->ne[2];
5219    //const int ne13 = src1->ne[3];
5220
5221    //const int ne0  = dst->ne[0];
5222    //const int ne1  = dst->ne[1];
5223    //const int ne2  = dst->ne[2];
5224    //const int ne3  = dst->ne[3];
5225    //const int ne   = ne0*ne1*ne2*ne3;
5226
5227    const int nb00 = src0->nb[0];
5228    const int nb01 = src0->nb[1];
5229    const int nb02 = src0->nb[2];
5230    //const int nb03 = src0->nb[3];
5231
5232    const int nb10 = src1->nb[0];
5233    const int nb11 = src1->nb[1];
5234    //const int nb12 = src1->nb[2];
5235    //const int nb13 = src1->nb[3];
5236
5237    //const int nb0  = dst->nb[0];
5238    const int nb1  = dst->nb[1];
5239    //const int nb2  = dst->nb[2];
5240    //const int nb3  = dst->nb[3];
5241
5242    const int ith = params->ith;
5243    const int nth = params->nth;
5244
5245    const int nk = ne00;
5246    const int nh = nk/2;
5247
5248    const int ew0 = ggml_up32(ne01);
5249
5250    GGML_ASSERT(ne00 % 2 == 1); // TODO: support even kernel sizes
5251    GGML_ASSERT(nb00 == sizeof(ggml_fp16_t));
5252    GGML_ASSERT(nb10 == sizeof(float));
5253
5254    if (params->type == GGML_TASK_INIT) {
5255        // TODO: fix this memset (wsize is overestimated)
5256        memset(params->wdata, 0, params->wsize);
5257
5258        // prepare kernel data (src0)
5259        {
5260            ggml_fp16_t * const wdata = (ggml_fp16_t *) params->wdata + 0;
5261
5262            for (int i02 = 0; i02 < ne02; i02++) {
5263                for (int i01 = 0; i01 < ne01; i01++) {
5264                    const ggml_fp16_t * const src = (ggml_fp16_t *)((char *) src0->data + i02*nb02 + i01*nb01);
5265                    ggml_fp16_t * dst_data = wdata + i02*ew0*ne00;
5266                    for (int i00 = 0; i00 < ne00; i00++) {
5267                        dst_data[i00*ew0 + i01] = src[i00];
5268                    }
5269                }
5270            }
5271        }
5272
5273        // prepare source data (src1)
5274        {
5275            ggml_fp16_t * const wdata = (ggml_fp16_t *) params->wdata + ne02*ew0*ne00;
5276
5277            for (int i11 = 0; i11 < ne11; i11++) {
5278                const float * const src = (float *)((char *) src1->data + i11*nb11);
5279                ggml_fp16_t * dst_data = wdata;
5280                for (int i10 = 0; i10 < ne10; i10++) {
5281                    dst_data[(i10 + nh)*ew0 + i11] = GGML_FP32_TO_FP16(src[i10]);
5282                }
5283            }
5284        }
5285
5286        return;
5287    }
5288
5289    if (params->type == GGML_TASK_FINALIZE) {
5290        return;
5291    }
5292
5293    // total rows in dst
5294    const int nr = ne02;
5295
5296    // rows per thread
5297    const int dr = (nr + nth - 1)/nth;
5298
5299    // row range for this thread
5300    const int ir0 = dr*ith;
5301    const int ir1 = MIN(ir0 + dr, nr);
5302
5303    for (int i1 = ir0; i1 < ir1; i1++) {
5304        float * dst_data = (float *)((char *) dst->data + i1*nb1);
5305        for (int i0 = 0; i0 < ne10; ++i0) {
5306            dst_data[i0] = 0;
5307            for (int k = -nh; k <= nh; k++) {
5308                float v = 0.0f;
5309                ggml_vec_dot_f16(ew0, &v,
5310                        (ggml_fp16_t *) params->wdata +   i1*ew0*ne00 +      (nh + k)*ew0,
5311                        (ggml_fp16_t *) params->wdata + ne02*ew0*ne00 + (i0 + nh + k)*ew0);
5312
5313                dst_data[i0] += v;
5314            }
5315        }
5316    }
5317}
5318
5319static void ggml_compute_forward_conv_1d_1s_f32(
5320        const struct ggml_compute_params * params,
5321        const struct ggml_tensor * src0,
5322        const struct ggml_tensor * src1,
5323              struct ggml_tensor * dst) {
5324    GGML_ASSERT(src0->type == GGML_TYPE_F32);
5325    GGML_ASSERT(src1->type == GGML_TYPE_F32);
5326    GGML_ASSERT( dst->type == GGML_TYPE_F32);
5327
5328    int64_t t0 = ggml_perf_time_us();
5329    UNUSED(t0);
5330
5331    const int ne00 = src0->ne[0];
5332    const int ne01 = src0->ne[1];
5333    const int ne02 = src0->ne[2];
5334    //const int ne03 = src0->ne[3];
5335
5336    const int ne10 = src1->ne[0];
5337    const int ne11 = src1->ne[1];
5338    //const int ne12 = src1->ne[2];
5339    //const int ne13 = src1->ne[3];
5340
5341    //const int ne0  = dst->ne[0];
5342    //const int ne1  = dst->ne[1];
5343    //const int ne2  = dst->ne[2];
5344    //const int ne3  = dst->ne[3];
5345    //const int ne   = ne0*ne1*ne2*ne3;
5346
5347    const int nb00 = src0->nb[0];
5348    const int nb01 = src0->nb[1];
5349    const int nb02 = src0->nb[2];
5350    //const int nb03 = src0->nb[3];
5351
5352    const int nb10 = src1->nb[0];
5353    const int nb11 = src1->nb[1];
5354    //const int nb12 = src1->nb[2];
5355    //const int nb13 = src1->nb[3];
5356
5357    //const int nb0  = dst->nb[0];
5358    const int nb1  = dst->nb[1];
5359    //const int nb2  = dst->nb[2];
5360    //const int nb3  = dst->nb[3];
5361
5362    const int ith = params->ith;
5363    const int nth = params->nth;
5364
5365    const int nk = ne00;
5366    const int nh = nk/2;
5367
5368    const int ew0 = ggml_up32(ne01);
5369
5370    GGML_ASSERT(ne00 % 2 == 1); // TODO: support even kernel sizes
5371    GGML_ASSERT(nb00 == sizeof(float));
5372    GGML_ASSERT(nb10 == sizeof(float));
5373
5374    if (params->type == GGML_TASK_INIT) {
5375        // TODO: fix this memset (wsize is overestimated)
5376        memset(params->wdata, 0, params->wsize);
5377
5378        // prepare kernel data (src0)
5379        {
5380            float * const wdata = (float *) params->wdata + 0;
5381
5382            for (int i02 = 0; i02 < ne02; i02++) {
5383                for (int i01 = 0; i01 < ne01; i01++) {
5384                    const float * const src = (float *)((char *) src0->data + i02*nb02 + i01*nb01);
5385                    float * dst_data = wdata + i02*ew0*ne00;
5386                    for (int i00 = 0; i00 < ne00; i00++) {
5387                        dst_data[i00*ew0 + i01] = src[i00];
5388                    }
5389                }
5390            }
5391        }
5392
5393        // prepare source data (src1)
5394        {
5395            float * const wdata = (float *) params->wdata + ne02*ew0*ne00;
5396
5397            for (int i11 = 0; i11 < ne11; i11++) {
5398                const float * const src = (float *)((char *) src1->data + i11*nb11);
5399                float * dst_data = wdata;
5400                for (int i10 = 0; i10 < ne10; i10++) {
5401                    dst_data[(i10 + nh)*ew0 + i11] = src[i10];
5402                }
5403            }
5404        }
5405
5406        return;
5407    }
5408
5409    if (params->type == GGML_TASK_FINALIZE) {
5410        return;
5411    }
5412
5413    // total rows in dst
5414    const int nr = ne02;
5415
5416    // rows per thread
5417    const int dr = (nr + nth - 1)/nth;
5418
5419    // row range for this thread
5420    const int ir0 = dr*ith;
5421    const int ir1 = MIN(ir0 + dr, nr);
5422
5423    for (int i1 = ir0; i1 < ir1; i1++) {
5424        float * dst_data = (float *)((char *) dst->data + i1*nb1);
5425        for (int i0 = 0; i0 < ne10; ++i0) {
5426            dst_data[i0] = 0;
5427            for (int k = -nh; k <= nh; k++) {
5428                float v = 0.0f;
5429                ggml_vec_dot_f32(ew0, &v,
5430                        (float *) params->wdata +   i1*ew0*ne00 +      (nh + k)*ew0,
5431                        (float *) params->wdata + ne02*ew0*ne00 + (i0 + nh + k)*ew0);
5432
5433                dst_data[i0] += v;
5434            }
5435        }
5436    }
5437}
5438
5439static void ggml_compute_forward_conv_1d_1s(
5440        const struct ggml_compute_params * params,
5441        const struct ggml_tensor * src0,
5442        const struct ggml_tensor * src1,
5443        struct ggml_tensor * dst) {
5444    switch (src0->type) {
5445        case GGML_TYPE_F16:
5446            {
5447                ggml_compute_forward_conv_1d_1s_f16_f32(params, src0, src1, dst);
5448            } break;
5449        case GGML_TYPE_F32:
5450            {
5451                ggml_compute_forward_conv_1d_1s_f32(params, src0, src1, dst);
5452            } break;
5453        case GGML_TYPE_I8:
5454        case GGML_TYPE_I16:
5455        case GGML_TYPE_I32:
5456        case GGML_TYPE_COUNT:
5457            {
5458                GGML_ASSERT(false);
5459            } break;
5460    }
5461}
5462
5463// ggml_compute_forward_conv_1d_2s
5464
5465static void ggml_compute_forward_conv_1d_2s_f16_f32(
5466        const struct ggml_compute_params * params,
5467        const struct ggml_tensor * src0,
5468        const struct ggml_tensor * src1,
5469              struct ggml_tensor * dst) {
5470    GGML_ASSERT(src0->type == GGML_TYPE_F16);
5471    GGML_ASSERT(src1->type == GGML_TYPE_F32);
5472    GGML_ASSERT( dst->type == GGML_TYPE_F32);
5473
5474    int64_t t0 = ggml_perf_time_us();
5475    UNUSED(t0);
5476
5477    const int ne00 = src0->ne[0];
5478    const int ne01 = src0->ne[1];
5479    const int ne02 = src0->ne[2];
5480    //const int ne03 = src0->ne[3];
5481
5482    const int ne10 = src1->ne[0];
5483    const int ne11 = src1->ne[1];
5484    //const int ne12 = src1->ne[2];
5485    //const int ne13 = src1->ne[3];
5486
5487    //const int ne0  = dst->ne[0];
5488    //const int ne1  = dst->ne[1];
5489    //const int ne2  = dst->ne[2];
5490    //const int ne3  = dst->ne[3];
5491    //const int ne   = ne0*ne1*ne2*ne3;
5492
5493    const int nb00 = src0->nb[0];
5494    const int nb01 = src0->nb[1];
5495    const int nb02 = src0->nb[2];
5496    //const int nb03 = src0->nb[3];
5497
5498    const int nb10 = src1->nb[0];
5499    const int nb11 = src1->nb[1];
5500    //const int nb12 = src1->nb[2];
5501    //const int nb13 = src1->nb[3];
5502
5503    //const int nb0  = dst->nb[0];
5504    const int nb1  = dst->nb[1];
5505    //const int nb2  = dst->nb[2];
5506    //const int nb3  = dst->nb[3];
5507
5508    const int ith = params->ith;
5509    const int nth = params->nth;
5510
5511    const int nk = ne00;
5512    const int nh = nk/2;
5513
5514    const int ew0 = ggml_up32(ne01);
5515
5516    GGML_ASSERT(ne00 % 2 == 1); // TODO: support even kernel sizes
5517    GGML_ASSERT(nb00 == sizeof(ggml_fp16_t));
5518    GGML_ASSERT(nb10 == sizeof(float));
5519
5520    if (params->type == GGML_TASK_INIT) {
5521        // TODO: fix this memset (wsize is overestimated)
5522        memset(params->wdata, 0, params->wsize);
5523
5524        // prepare kernel data (src0)
5525        {
5526            ggml_fp16_t * const wdata = (ggml_fp16_t *) params->wdata + 0;
5527
5528            for (int i02 = 0; i02 < ne02; i02++) {
5529                for (int i01 = 0; i01 < ne01; i01++) {
5530                    const ggml_fp16_t * const src = (ggml_fp16_t *)((char *) src0->data + i02*nb02 + i01*nb01);
5531                    ggml_fp16_t * dst_data = wdata + i02*ew0*ne00;
5532                    for (int i00 = 0; i00 < ne00; i00++) {
5533                        dst_data[i00*ew0 + i01] = src[i00];
5534                    }
5535                }
5536            }
5537        }
5538
5539        // prepare source data (src1)
5540        {
5541            ggml_fp16_t * const wdata = (ggml_fp16_t *) params->wdata + ne02*ew0*ne00;
5542
5543            for (int i11 = 0; i11 < ne11; i11++) {
5544                const float * const src = (float *)((char *) src1->data + i11*nb11);
5545                ggml_fp16_t * dst_data = wdata;
5546                for (int i10 = 0; i10 < ne10; i10++) {
5547                    dst_data[(i10 + nh)*ew0 + i11] = GGML_FP32_TO_FP16(src[i10]);
5548                }
5549            }
5550        }
5551
5552        return;
5553    }
5554
5555    if (params->type == GGML_TASK_FINALIZE) {
5556        return;
5557    }
5558
5559    // total rows in dst
5560    const int nr = ne02;
5561
5562    // rows per thread
5563    const int dr = (nr + nth - 1)/nth;
5564
5565    // row range for this thread
5566    const int ir0 = dr*ith;
5567    const int ir1 = MIN(ir0 + dr, nr);
5568
5569    for (int i1 = ir0; i1 < ir1; i1++) {
5570        float * dst_data = (float *)((char *) dst->data + i1*nb1);
5571        for (int i0 = 0; i0 < ne10; i0 += 2) {
5572            dst_data[i0/2] = 0;
5573            for (int k = -nh; k <= nh; k++) {
5574                float v = 0.0f;
5575                ggml_vec_dot_f16(ew0, &v,
5576                        (ggml_fp16_t *) params->wdata +   i1*ew0*ne00 +      (nh + k)*ew0,
5577                        (ggml_fp16_t *) params->wdata + ne02*ew0*ne00 + (i0 + nh + k)*ew0);
5578
5579                dst_data[i0/2] += v;
5580            }
5581        }
5582    }
5583}
5584
5585static void ggml_compute_forward_conv_1d_2s_f32(
5586        const struct ggml_compute_params * params,
5587        const struct ggml_tensor * src0,
5588        const struct ggml_tensor * src1,
5589              struct ggml_tensor * dst) {
5590    GGML_ASSERT(src0->type == GGML_TYPE_F32);
5591    GGML_ASSERT(src1->type == GGML_TYPE_F32);
5592    GGML_ASSERT( dst->type == GGML_TYPE_F32);
5593
5594    int64_t t0 = ggml_perf_time_us();
5595    UNUSED(t0);
5596
5597    const int ne00 = src0->ne[0];
5598    const int ne01 = src0->ne[1];
5599    const int ne02 = src0->ne[2];
5600    //const int ne03 = src0->ne[3];
5601
5602    const int ne10 = src1->ne[0];
5603    const int ne11 = src1->ne[1];
5604    //const int ne12 = src1->ne[2];
5605    //const int ne13 = src1->ne[3];
5606
5607    //const int ne0  = dst->ne[0];
5608    //const int ne1  = dst->ne[1];
5609    //const int ne2  = dst->ne[2];
5610    //const int ne3  = dst->ne[3];
5611    //const int ne   = ne0*ne1*ne2*ne3;
5612
5613    const int nb00 = src0->nb[0];
5614    const int nb01 = src0->nb[1];
5615    const int nb02 = src0->nb[2];
5616    //const int nb03 = src0->nb[3];
5617
5618    const int nb10 = src1->nb[0];
5619    const int nb11 = src1->nb[1];
5620    //const int nb12 = src1->nb[2];
5621    //const int nb13 = src1->nb[3];
5622
5623    //const int nb0  = dst->nb[0];
5624    const int nb1  = dst->nb[1];
5625    //const int nb2  = dst->nb[2];
5626    //const int nb3  = dst->nb[3];
5627
5628    const int ith = params->ith;
5629    const int nth = params->nth;
5630
5631    const int nk = ne00;
5632    const int nh = nk/2;
5633
5634    const int ew0 = ggml_up32(ne01);
5635
5636    GGML_ASSERT(ne00 % 2 == 1); // TODO: support even kernel sizes
5637    GGML_ASSERT(nb00 == sizeof(float));
5638    GGML_ASSERT(nb10 == sizeof(float));
5639
5640    if (params->type == GGML_TASK_INIT) {
5641        // TODO: fix this memset (wsize is overestimated)
5642        memset(params->wdata, 0, params->wsize);
5643
5644        // prepare kernel data (src0)
5645        {
5646            float * const wdata = (float *) params->wdata + 0;
5647
5648            for (int i02 = 0; i02 < ne02; i02++) {
5649                for (int i01 = 0; i01 < ne01; i01++) {
5650                    const float * const src = (float *)((char *) src0->data + i02*nb02 + i01*nb01);
5651                    float * dst_data = wdata + i02*ew0*ne00;
5652                    for (int i00 = 0; i00 < ne00; i00++) {
5653                        dst_data[i00*ew0 + i01] = src[i00];
5654                    }
5655                }
5656            }
5657        }
5658
5659        // prepare source data (src1)
5660        {
5661            float * const wdata = (float *) params->wdata + ne02*ew0*ne00;
5662
5663            for (int i11 = 0; i11 < ne11; i11++) {
5664                const float * const src = (float *)((char *) src1->data + i11*nb11);
5665                float * dst_data = wdata;
5666                for (int i10 = 0; i10 < ne10; i10++) {
5667                    dst_data[(i10 + nh)*ew0 + i11] = src[i10];
5668                }
5669            }
5670        }
5671
5672        return;
5673    }
5674
5675    if (params->type == GGML_TASK_FINALIZE) {
5676        return;
5677    }
5678
5679    // total rows in dst
5680    const int nr = ne02;
5681
5682    // rows per thread
5683    const int dr = (nr + nth - 1)/nth;
5684
5685    // row range for this thread
5686    const int ir0 = dr*ith;
5687    const int ir1 = MIN(ir0 + dr, nr);
5688
5689    for (int i1 = ir0; i1 < ir1; i1++) {
5690        float * dst_data = (float *)((char *) dst->data + i1*nb1);
5691        for (int i0 = 0; i0 < ne10; i0 += 2) {
5692            dst_data[i0/2] = 0;
5693            for (int k = -nh; k <= nh; k++) {
5694                float v = 0.0f;
5695                ggml_vec_dot_f32(ew0, &v,
5696                        (float *) params->wdata +   i1*ew0*ne00 +      (nh + k)*ew0,
5697                        (float *) params->wdata + ne02*ew0*ne00 + (i0 + nh + k)*ew0);
5698
5699                dst_data[i0/2] += v;
5700            }
5701        }
5702    }
5703}
5704
5705static void ggml_compute_forward_conv_1d_2s(
5706        const struct ggml_compute_params * params,
5707        const struct ggml_tensor * src0,
5708        const struct ggml_tensor * src1,
5709        struct ggml_tensor * dst) {
5710    switch (src0->type) {
5711        case GGML_TYPE_F16:
5712            {
5713                ggml_compute_forward_conv_1d_2s_f16_f32(params, src0, src1, dst);
5714            } break;
5715        case GGML_TYPE_F32:
5716            {
5717                ggml_compute_forward_conv_1d_2s_f32(params, src0, src1, dst);
5718            } break;
5719        case GGML_TYPE_I8:
5720        case GGML_TYPE_I16:
5721        case GGML_TYPE_I32:
5722        case GGML_TYPE_COUNT:
5723            {
5724                GGML_ASSERT(false);
5725            } break;
5726    }
5727}
5728
5729// ggml_compute_forward_flash_attn
5730
5731static void ggml_compute_forward_flash_attn_f32(
5732        const struct ggml_compute_params * params,
5733        const struct ggml_tensor * q,
5734        const struct ggml_tensor * k,
5735        const struct ggml_tensor * v,
5736        const bool masked,
5737             struct ggml_tensor * dst) {
5738    int64_t t0 = ggml_perf_time_us();
5739    UNUSED(t0);
5740
5741    const int neq0 = q->ne[0];
5742    const int neq1 = q->ne[1];
5743    const int neq2 = q->ne[2];
5744    const int neq3 = q->ne[3];
5745
5746    const int nek0 = k->ne[0];
5747    const int nek1 = k->ne[1];
5748    //const int nek2 = k->ne[2];
5749    //const int nek3 = k->ne[3];
5750
5751    //const int nev0 = v->ne[0];
5752    const int nev1 = v->ne[1];
5753    //const int nev2 = v->ne[2];
5754    //const int nev3 = v->ne[3];
5755
5756    const int ne0  = dst->ne[0];
5757    const int ne1  = dst->ne[1];
5758    //const int ne2  = dst->ne[2];
5759    //const int ne3  = dst->ne[3];
5760
5761    const int nbk0 = k->nb[0];
5762    const int nbk1 = k->nb[1];
5763    const int nbk2 = k->nb[2];
5764    const int nbk3 = k->nb[3];
5765
5766    const int nbq0 = q->nb[0];
5767    const int nbq1 = q->nb[1];
5768    const int nbq2 = q->nb[2];
5769    const int nbq3 = q->nb[3];
5770
5771    const int nbv0 = v->nb[0];
5772    const int nbv1 = v->nb[1];
5773    const int nbv2 = v->nb[2];
5774    const int nbv3 = v->nb[3];
5775
5776    const int nb0  = dst->nb[0];
5777    const int nb1  = dst->nb[1];
5778    const int nb2  = dst->nb[2];
5779    const int nb3  = dst->nb[3];
5780
5781    const int ith = params->ith;
5782    const int nth = params->nth;
5783
5784    const int D = neq0;
5785    const int N = neq1;
5786    const int P = nek1 - N;
5787    const int M = P + N;
5788
5789    GGML_ASSERT(ne0 == D);
5790    GGML_ASSERT(ne1 == N);
5791    GGML_ASSERT(P >= 0);
5792
5793    GGML_ASSERT(nbq0 == sizeof(float));
5794    GGML_ASSERT(nbk0 == sizeof(float));
5795    GGML_ASSERT(nbv0 == sizeof(float));
5796
5797    GGML_ASSERT(neq0 == D);
5798    GGML_ASSERT(nek0 == D);
5799    GGML_ASSERT(nev1 == D);
5800
5801    GGML_ASSERT(neq1 == N);
5802    GGML_ASSERT(nek1 == N + P);
5803    GGML_ASSERT(nev1 == D);
5804
5805    // dst cannot be transposed or permuted
5806    GGML_ASSERT(nb0 == sizeof(float));
5807    GGML_ASSERT(nb0 <= nb1);
5808    GGML_ASSERT(nb1 <= nb2);
5809    GGML_ASSERT(nb2 <= nb3);
5810
5811    if (params->type == GGML_TASK_INIT) {
5812        return;
5813    }
5814
5815    if (params->type == GGML_TASK_FINALIZE) {
5816        return;
5817    }
5818
5819    // parallelize by q rows using ggml_vec_dot_f32
5820
5821    // total rows in q
5822    const int nr = neq1*neq2*neq3;
5823
5824    // rows per thread
5825    const int dr = (nr + nth - 1)/nth;
5826
5827    // row range for this thread
5828    const int ir0 = dr*ith;
5829    const int ir1 = MIN(ir0 + dr, nr);
5830
5831    const float scale = 1.0/sqrt((double) D);
5832
5833    //printf("P=%d N=%d D=%d ir0=%d ir1=%d scale = %f\n", P, N, D, ir0, ir1, scale);
5834
5835    for (int ir = ir0; ir < ir1; ++ir) {
5836        // q indices
5837        const int iq3 = ir/(neq2*neq1);
5838        const int iq2 = (ir - iq3*neq2*neq1)/neq1;
5839        const int iq1 = (ir - iq3*neq2*neq1 - iq2*neq1);
5840
5841        float * S = (float *) params->wdata + ith*(M + CACHE_LINE_SIZE_F32);
5842
5843        for (int ic = 0; ic < nek1; ++ic) {
5844            // k indices
5845            const int ik3 = iq3;
5846            const int ik2 = iq2;
5847            const int ik1 = ic;
5848
5849            // S indices
5850            const int i1 = ik1;
5851
5852            ggml_vec_dot_f32(neq0,
5853                    S + i1,
5854                    (float *) ((char *) k->data + (ik1*nbk1 + ik2*nbk2 + ik3*nbk3)),
5855                    (float *) ((char *) q->data + (iq1*nbq1 + iq2*nbq2 + iq3*nbq3)));
5856        }
5857
5858        // scale
5859        ggml_vec_scale_f32(nek1, S, scale);
5860
5861        if (masked) {
5862            for (int i = P; i < M; i++) {
5863                if (i > P + iq1) {
5864                    S[i] = -INFINITY;
5865                }
5866            }
5867        }
5868
5869        // softmax
5870        {
5871            float max = -INFINITY;
5872            for (int i = 0; i < M; i++) {
5873                max = MAX(max, S[i]);
5874            }
5875
5876            ggml_float sum = 0.0;
5877
5878            uint16_t ss;
5879            for (int i = 0; i < M; i++) {
5880                if (S[i] == -INFINITY) {
5881                    S[i] = 0.0;
5882                } else {
5883                    //const float val = (S[i] == -INFINITY) ? 0.0 : exp(S[i] - max);
5884                    ggml_fp16_t s = GGML_FP32_TO_FP16(S[i] - max);
5885                    memcpy(&ss, &s, sizeof(ss));
5886                    const float val = GGML_FP16_TO_FP32(table_exp_f16[ss]);
5887                    sum += val;
5888                    S[i] = val;
5889                }
5890            }
5891
5892            assert(sum > 0.0f);
5893
5894            sum = 1.0/sum;
5895            ggml_vec_scale_f32(M, S, sum);
5896        }
5897
5898        for (int ic = 0; ic < nev1; ++ic) {
5899            // dst indices
5900            const int i1 = iq1;
5901            const int i2 = iq2;
5902            const int i3 = iq3;
5903
5904            ggml_vec_dot_f32(nek1,
5905                    (float *) ((char *) dst->data + (ic*nb0 + i1*nb1  + i2*nb2  + i3*nb3)),
5906                    (float *) ((char *) v->data   + (         ic*nbv1 + i2*nbv2 + i3*nbv3)),
5907                    S);
5908        }
5909    }
5910}
5911
5912static void ggml_compute_forward_flash_attn_f16(
5913        const struct ggml_compute_params * params,
5914        const struct ggml_tensor * q,
5915        const struct ggml_tensor * k,
5916        const struct ggml_tensor * v,
5917        const bool masked,
5918             struct ggml_tensor * dst) {
5919    int64_t t0 = ggml_perf_time_us();
5920    UNUSED(t0);
5921
5922    const int neq0 = q->ne[0];
5923    const int neq1 = q->ne[1];
5924    const int neq2 = q->ne[2];
5925    const int neq3 = q->ne[3];
5926
5927    const int nek0 = k->ne[0];
5928    const int nek1 = k->ne[1];
5929    //const int nek2 = k->ne[2];
5930    //const int nek3 = k->ne[3];
5931
5932    //const int nev0 = v->ne[0];
5933    const int nev1 = v->ne[1];
5934    //const int nev2 = v->ne[2];
5935    //const int nev3 = v->ne[3];
5936
5937    const int ne0  = dst->ne[0];
5938    const int ne1  = dst->ne[1];
5939    //const int ne2  = dst->ne[2];
5940    //const int ne3  = dst->ne[3];
5941
5942    const int nbk0 = k->nb[0];
5943    const int nbk1 = k->nb[1];
5944    const int nbk2 = k->nb[2];
5945    const int nbk3 = k->nb[3];
5946
5947    const int nbq0 = q->nb[0];
5948    const int nbq1 = q->nb[1];
5949    const int nbq2 = q->nb[2];
5950    const int nbq3 = q->nb[3];
5951
5952    const int nbv0 = v->nb[0];
5953    const int nbv1 = v->nb[1];
5954    const int nbv2 = v->nb[2];
5955    const int nbv3 = v->nb[3];
5956
5957    const int nb0  = dst->nb[0];
5958    const int nb1  = dst->nb[1];
5959    const int nb2  = dst->nb[2];
5960    const int nb3  = dst->nb[3];
5961
5962    const int ith = params->ith;
5963    const int nth = params->nth;
5964
5965    const int D = neq0;
5966    const int N = neq1;
5967    const int P = nek1 - N;
5968    const int M = P + N;
5969
5970    GGML_ASSERT(ne0 == D);
5971    GGML_ASSERT(ne1 == N);
5972    GGML_ASSERT(P >= 0);
5973
5974    GGML_ASSERT(nbq0 == sizeof(ggml_fp16_t));
5975    GGML_ASSERT(nbk0 == sizeof(ggml_fp16_t));
5976    GGML_ASSERT(nbv0 == sizeof(ggml_fp16_t));
5977
5978    GGML_ASSERT(neq0 == D);
5979    GGML_ASSERT(nek0 == D);
5980    GGML_ASSERT(nev1 == D);
5981
5982    GGML_ASSERT(neq1 == N);
5983    GGML_ASSERT(nek1 == N + P);
5984    GGML_ASSERT(nev1 == D);
5985
5986    // dst cannot be transposed or permuted
5987    GGML_ASSERT(nb0 == sizeof(float));
5988    GGML_ASSERT(nb0 <= nb1);
5989    GGML_ASSERT(nb1 <= nb2);
5990    GGML_ASSERT(nb2 <= nb3);
5991
5992    if (params->type == GGML_TASK_INIT) {
5993        return;
5994    }
5995
5996    if (params->type == GGML_TASK_FINALIZE) {
5997        return;
5998    }
5999
6000    // parallelize by q rows using ggml_vec_dot_f32
6001
6002    // total rows in q
6003    const int nr = neq1*neq2*neq3;
6004
6005    // rows per thread
6006    const int dr = (nr + nth - 1)/nth;
6007
6008    // row range for this thread
6009    const int ir0 = dr*ith;
6010    const int ir1 = MIN(ir0 + dr, nr);
6011
6012    const float scale = 1.0/sqrt((double) D);
6013
6014    //printf("P=%d N=%d D=%d ir0=%d ir1=%d scale = %f\n", P, N, D, ir0, ir1, scale);
6015
6016    for (int ir = ir0; ir < ir1; ++ir) {
6017        // q indices
6018        const int iq3 = ir/(neq2*neq1);
6019        const int iq2 = (ir - iq3*neq2*neq1)/neq1;
6020        const int iq1 = (ir - iq3*neq2*neq1 - iq2*neq1);
6021
6022        float * S = (float *) params->wdata + ith*(2*M + CACHE_LINE_SIZE_F32);
6023
6024        for (int ic = 0; ic < nek1; ++ic) {
6025            // k indices
6026            const int ik3 = iq3;
6027            const int ik2 = iq2;
6028            const int ik1 = ic;
6029
6030            // S indices
6031            const int i1 = ik1;
6032
6033            ggml_vec_dot_f16(neq0,
6034                    S + i1,
6035                    (ggml_fp16_t *) ((char *) k->data + (ik1*nbk1 + ik2*nbk2 + ik3*nbk3)),
6036                    (ggml_fp16_t *) ((char *) q->data + (iq1*nbq1 + iq2*nbq2 + iq3*nbq3)));
6037        }
6038
6039        // scale
6040        ggml_vec_scale_f32(nek1, S, scale);
6041
6042        if (masked) {
6043            for (int i = P; i < M; i++) {
6044                if (i > P + iq1) {
6045                    S[i] = -INFINITY;
6046                }
6047            }
6048        }
6049
6050        // softmax
6051        {
6052            float max = -INFINITY;
6053            for (int i = 0; i < M; i++) {
6054                max = MAX(max, S[i]);
6055            }
6056
6057            ggml_float sum = 0.0;
6058
6059            uint16_t ss;
6060            for (int i = 0; i < M; i++) {
6061                if (S[i] == -INFINITY) {
6062                    S[i] = 0.0;
6063                } else {
6064                    //const float val = (S[i] == -INFINITY) ? 0.0 : exp(S[i] - max);
6065                    ggml_fp16_t s = GGML_FP32_TO_FP16(S[i] - max);
6066                    memcpy(&ss, &s, sizeof(ss));
6067                    const float val = GGML_FP16_TO_FP32(table_exp_f16[ss]);
6068                    sum += val;
6069                    S[i] = val;
6070                }
6071            }
6072
6073            assert(sum > 0.0f);
6074
6075            sum = 1.0/sum;
6076            ggml_vec_scale_f32(M, S, sum);
6077        }
6078
6079        ggml_fp16_t * S16 = (ggml_fp16_t *) ((float *) params->wdata + ith*(2*M + CACHE_LINE_SIZE_F32) + M);
6080
6081        for (int i = 0; i < M; i++) {
6082            S16[i] = GGML_FP32_TO_FP16(S[i]);
6083        }
6084
6085        for (int ic = 0; ic < nev1; ++ic) {
6086            // dst indices
6087            const int i1 = iq1;
6088            const int i2 = iq2;
6089            const int i3 = iq3;
6090
6091            ggml_vec_dot_f16(nek1,
6092                    (float *)       ((char *) dst->data + (ic*nb0 + i1*nb1  + i2*nb2  + i3*nb3)),
6093                    (ggml_fp16_t *) ((char *) v->data   + (         ic*nbv1 + i2*nbv2 + i3*nbv3)),
6094                    S16);
6095        }
6096    }
6097}
6098
6099static void ggml_compute_forward_flash_attn(
6100        const struct ggml_compute_params * params,
6101        const struct ggml_tensor * q,
6102        const struct ggml_tensor * k,
6103        const struct ggml_tensor * v,
6104        const bool masked,
6105        struct ggml_tensor * dst) {
6106    switch (q->type) {
6107        case GGML_TYPE_F16:
6108            {
6109                ggml_compute_forward_flash_attn_f16(params, q, k, v, masked, dst);
6110            } break;
6111        case GGML_TYPE_F32:
6112            {
6113                ggml_compute_forward_flash_attn_f32(params, q, k, v, masked, dst);
6114            } break;
6115        case GGML_TYPE_I8:
6116        case GGML_TYPE_I16:
6117        case GGML_TYPE_I32:
6118        case GGML_TYPE_COUNT:
6119            {
6120                assert(false);
6121            } break;
6122    }
6123}
6124
6125// ggml_compute_forward_flash_ff
6126
6127static void ggml_compute_forward_flash_ff_f16(
6128        const struct ggml_compute_params * params,
6129        const struct ggml_tensor * a,  // F16
6130        const struct ggml_tensor * b0, // F16 fc_w
6131        const struct ggml_tensor * b1, // F32 fc_b
6132        const struct ggml_tensor * c0, // F16 proj_w
6133        const struct ggml_tensor * c1, // F32 proj_b
6134        struct ggml_tensor * dst) {
6135    int64_t t0 = ggml_perf_time_us();
6136    UNUSED(t0);
6137
6138    const int nea0 = a->ne[0];
6139    const int nea1 = a->ne[1];
6140    const int nea2 = a->ne[2];
6141    const int nea3 = a->ne[3];
6142
6143    const int neb00 = b0->ne[0];
6144    const int neb01 = b0->ne[1];
6145    //const int neb02 = b0->ne[2];
6146    //const int neb03 = b0->ne[3];
6147
6148    const int neb10 = b1->ne[0];
6149    const int neb11 = b1->ne[1];
6150    //const int neb12 = b1->ne[2];
6151    //const int neb13 = b1->ne[3];
6152
6153    const int nec00 = c0->ne[0];
6154    const int nec01 = c0->ne[1];
6155    //const int nec02 = c0->ne[2];
6156    //const int nec03 = c0->ne[3];
6157
6158    const int nec10 = c1->ne[0];
6159    const int nec11 = c1->ne[1];
6160    //const int nec12 = c1->ne[2];
6161    //const int nec13 = c1->ne[3];
6162
6163    const int ne0 = dst->ne[0];
6164    const int ne1 = dst->ne[1];
6165    const int ne2 = dst->ne[2];
6166    //const int ne3 = dst->ne[3];
6167
6168    const int nba0 = a->nb[0];
6169    const int nba1 = a->nb[1];
6170    const int nba2 = a->nb[2];
6171    const int nba3 = a->nb[3];
6172
6173    const int nbb00 = b0->nb[0];
6174    const int nbb01 = b0->nb[1];
6175    const int nbb02 = b0->nb[2];
6176    const int nbb03 = b0->nb[3];
6177
6178    const int nbb10 = b1->nb[0];
6179    //const int nbb11 = b1->nb[1];
6180    //const int nbb12 = b1->nb[2];
6181    //const int nbb13 = b1->nb[3];
6182
6183    const int nbc00 = c0->nb[0];
6184    const int nbc01 = c0->nb[1];
6185    const int nbc02 = c0->nb[2];
6186    const int nbc03 = c0->nb[3];
6187
6188    const int nbc10 = c1->nb[0];
6189    //const int nbc11 = c1->nb[1];
6190    //const int nbc12 = c1->nb[2];
6191    //const int nbc13 = c1->nb[3];
6192
6193    const int nb0 = dst->nb[0];
6194    const int nb1 = dst->nb[1];
6195    const int nb2 = dst->nb[2];
6196    const int nb3 = dst->nb[3];
6197
6198    const int ith = params->ith;
6199    const int nth = params->nth;
6200
6201    const int D = nea0;
6202    //const int N = nea1;
6203    const int M = neb01;
6204
6205    GGML_ASSERT(ne0 == nea0);
6206    GGML_ASSERT(ne1 == nea1);
6207    GGML_ASSERT(ne2 == nea2);
6208
6209    GGML_ASSERT(nba0  == sizeof(ggml_fp16_t));
6210    GGML_ASSERT(nbb00 == sizeof(ggml_fp16_t));
6211    GGML_ASSERT(nbb10 == sizeof(float));
6212    GGML_ASSERT(nbc00 == sizeof(ggml_fp16_t));
6213    GGML_ASSERT(nbc10 == sizeof(float));
6214
6215    GGML_ASSERT(neb00 == D);
6216    GGML_ASSERT(neb01 == M);
6217    GGML_ASSERT(neb10 == M);
6218    GGML_ASSERT(neb11 == 1);
6219
6220    GGML_ASSERT(nec00 == M);
6221    GGML_ASSERT(nec01 == D);
6222    GGML_ASSERT(nec10 == D);
6223    GGML_ASSERT(nec11 == 1);
6224
6225    // dst cannot be transposed or permuted
6226    GGML_ASSERT(nb0 == sizeof(float));
6227    GGML_ASSERT(nb0 <= nb1);
6228    GGML_ASSERT(nb1 <= nb2);
6229    GGML_ASSERT(nb2 <= nb3);
6230
6231    if (params->type == GGML_TASK_INIT) {
6232        return;
6233    }
6234
6235    if (params->type == GGML_TASK_FINALIZE) {
6236        return;
6237    }
6238
6239    // parallelize by a rows using ggml_vec_dot_f32
6240
6241    // total rows in a
6242    const int nr = nea1*nea2*nea3;
6243
6244    // rows per thread
6245    const int dr = (nr + nth - 1)/nth;
6246
6247    // row range for this thread
6248    const int ir0 = dr*ith;
6249    const int ir1 = MIN(ir0 + dr, nr);
6250
6251    for (int ir = ir0; ir < ir1; ++ir) {
6252        // a indices
6253        const int ia3 = ir/(nea2*nea1);
6254        const int ia2 = (ir - ia3*nea2*nea1)/nea1;
6255        const int ia1 = (ir - ia3*nea2*nea1 - ia2*nea1);
6256
6257        float * S = (float *) params->wdata + ith*(2*M + CACHE_LINE_SIZE_F32);
6258
6259        for (int ic = 0; ic < neb01; ++ic) {
6260            // b0 indices
6261            const int ib03 = ia3;
6262            const int ib02 = ia2;
6263            const int ib01 = ic;
6264
6265            // S indices
6266            const int i1 = ib01;
6267
6268            ggml_vec_dot_f16(nea0,
6269                    S + i1,
6270                    (ggml_fp16_t *) ((char *) b0->data + (ib01*nbb01 + ib02*nbb02 + ib03*nbb03)),
6271                    (ggml_fp16_t *) ((char *)  a->data + ( ia1*nba1  +  ia2*nba2  +  ia3*nba3)));
6272        }
6273
6274        ggml_vec_add_f32(neb01, S, S, (float *) b1->data);
6275        //ggml_vec_gelu_f32(neb01, S, S);
6276
6277        ggml_fp16_t * S16 = (ggml_fp16_t *) ((float *) params->wdata + ith*(2*M + CACHE_LINE_SIZE_F32) + M);
6278
6279        for (int i = 0; i < M; i++) {
6280            S16[i] = GGML_FP32_TO_FP16(S[i]);
6281        }
6282
6283        ggml_vec_gelu_f16(neb01, S16, S16);
6284
6285        {
6286            // dst indices
6287            const int i1 = ia1;
6288            const int i2 = ia2;
6289            const int i3 = ia3;
6290
6291            for (int ic = 0; ic < nec01; ++ic) {
6292
6293                ggml_vec_dot_f16(neb01,
6294                        (float *)       ((char *) dst->data + (ic*nb0 + i1*nb1   + i2*nb2   + i3*nb3)),
6295                        (ggml_fp16_t *) ((char *) c0->data  + (         ic*nbc01 + i2*nbc02 + i3*nbc03)),
6296                        S16);
6297            }
6298
6299            ggml_vec_add_f32(nec01,
6300                    (float *) ((char *) dst->data + (i1*nb1 + i2*nb2 + i3*nb3)),
6301                    (float *) ((char *) dst->data + (i1*nb1 + i2*nb2 + i3*nb3)),
6302                    (float *) c1->data);
6303        }
6304    }
6305}
6306
6307static void ggml_compute_forward_flash_ff(
6308        const struct ggml_compute_params * params,
6309        const struct ggml_tensor * a,
6310        const struct ggml_tensor * b0,
6311        const struct ggml_tensor * b1,
6312        const struct ggml_tensor * c0,
6313        const struct ggml_tensor * c1,
6314        struct ggml_tensor * dst) {
6315    switch (b0->type) {
6316        case GGML_TYPE_F16:
6317            {
6318                ggml_compute_forward_flash_ff_f16(params, a, b0, b1, c0, c1, dst);
6319            } break;
6320        case GGML_TYPE_F32:
6321            {
6322                GGML_ASSERT(false); // TODO
6323            } break;
6324        case GGML_TYPE_I8:
6325        case GGML_TYPE_I16:
6326        case GGML_TYPE_I32:
6327        case GGML_TYPE_COUNT:
6328            {
6329                assert(false);
6330            } break;
6331    }
6332}
6333
6334/////////////////////////////////
6335
6336static void ggml_compute_forward(struct ggml_compute_params * params, struct ggml_tensor * tensor) {
6337    assert(params);
6338
6339    switch (tensor->op) {
6340        case GGML_OP_DUP:
6341            {
6342                ggml_compute_forward_dup(params, tensor->src0, tensor);
6343            } break;
6344        case GGML_OP_ADD:
6345            {
6346                ggml_compute_forward_add(params, tensor->src0, tensor->src1, tensor);
6347            } break;
6348        case GGML_OP_SUB:
6349            {
6350                ggml_compute_forward_sub(params, tensor->src0, tensor->src1, tensor);
6351            } break;
6352        case GGML_OP_MUL:
6353            {
6354                ggml_compute_forward_mul(params, tensor->src0, tensor->src1, tensor);
6355            } break;
6356        case GGML_OP_DIV:
6357            {
6358                ggml_compute_forward_div(params, tensor->src0, tensor->src1, tensor);
6359            } break;
6360        case GGML_OP_SQR:
6361            {
6362                ggml_compute_forward_sqr(params, tensor->src0, tensor);
6363            } break;
6364        case GGML_OP_SQRT:
6365            {
6366                ggml_compute_forward_sqrt(params, tensor->src0, tensor);
6367            } break;
6368        case GGML_OP_SUM:
6369            {
6370                ggml_compute_forward_sum(params, tensor->src0, tensor);
6371            } break;
6372        case GGML_OP_MEAN:
6373            {
6374                ggml_compute_forward_mean(params, tensor->src0, tensor);
6375            } break;
6376        case GGML_OP_REPEAT:
6377            {
6378                ggml_compute_forward_repeat(params, tensor->src0, tensor);
6379            } break;
6380        case GGML_OP_ABS:
6381            {
6382                ggml_compute_forward_abs(params, tensor->src0, tensor);
6383            } break;
6384        case GGML_OP_SGN:
6385            {
6386                ggml_compute_forward_sgn(params, tensor->src0, tensor);
6387            } break;
6388        case GGML_OP_NEG:
6389            {
6390                ggml_compute_forward_neg(params, tensor->src0, tensor);
6391            } break;
6392        case GGML_OP_STEP:
6393            {
6394                ggml_compute_forward_step(params, tensor->src0, tensor);
6395            } break;
6396        case GGML_OP_RELU:
6397            {
6398                ggml_compute_forward_relu(params, tensor->src0, tensor);
6399            } break;
6400        case GGML_OP_GELU:
6401            {
6402                ggml_compute_forward_gelu(params, tensor->src0, tensor);
6403            } break;
6404        case GGML_OP_NORM:
6405            {
6406                ggml_compute_forward_norm(params, tensor->src0, tensor);
6407            } break;
6408        case GGML_OP_MUL_MAT:
6409            {
6410                ggml_compute_forward_mul_mat(params, tensor->src0, tensor->src1, tensor);
6411            } break;
6412        case GGML_OP_SCALE:
6413            {
6414                ggml_compute_forward_scale(params, tensor->src0, tensor->src1, tensor);
6415            } break;
6416        case GGML_OP_CPY:
6417            {
6418                ggml_compute_forward_cpy(params, tensor->src0, tensor);
6419            } break;
6420        case GGML_OP_RESHAPE:
6421            {
6422                ggml_compute_forward_reshape(params, tensor->src0, tensor);
6423            } break;
6424        case GGML_OP_VIEW:
6425            {
6426                ggml_compute_forward_view(params, tensor->src0);
6427            } break;
6428        case GGML_OP_PERMUTE:
6429            {
6430                ggml_compute_forward_permute(params, tensor->src0);
6431            } break;
6432        case GGML_OP_TRANSPOSE:
6433            {
6434                ggml_compute_forward_transpose(params, tensor->src0);
6435            } break;
6436        case GGML_OP_GET_ROWS:
6437            {
6438                ggml_compute_forward_get_rows(params, tensor->src0, tensor->src1, tensor);
6439            } break;
6440        case GGML_OP_DIAG_MASK_INF:
6441            {
6442                ggml_compute_forward_diag_mask_inf(params, tensor->src0, tensor->src1, tensor);
6443            } break;
6444        case GGML_OP_SOFT_MAX:
6445            {
6446                ggml_compute_forward_soft_max(params, tensor->src0, tensor);
6447            } break;
6448        case GGML_OP_ROPE:
6449            {
6450                ggml_compute_forward_rope(params, tensor->src0, tensor->src1, tensor);
6451            } break;
6452        case GGML_OP_CONV_1D_1S:
6453            {
6454                ggml_compute_forward_conv_1d_1s(params, tensor->src0, tensor->src1, tensor);
6455            } break;
6456        case GGML_OP_CONV_1D_2S:
6457            {
6458                ggml_compute_forward_conv_1d_2s(params, tensor->src0, tensor->src1, tensor);
6459            } break;
6460        case GGML_OP_FLASH_ATTN:
6461            {
6462                int32_t t = ggml_get_i32_1d(tensor->opt[1], 0);
6463                GGML_ASSERT(t == 0 || t == 1);
6464                bool masked = t != 0;
6465                ggml_compute_forward_flash_attn(params, tensor->src0, tensor->src1, tensor->opt[0], masked, tensor);
6466            } break;
6467        case GGML_OP_FLASH_FF:
6468            {
6469                ggml_compute_forward_flash_ff(params, tensor->src0, tensor->src1, tensor->opt[0], tensor->opt[1], tensor->opt[2], tensor);
6470            } break;
6471        case GGML_OP_NONE:
6472            {
6473                // nop
6474            } break;
6475        case GGML_OP_COUNT:
6476            {
6477                GGML_ASSERT(false);
6478            } break;
6479    }
6480}
6481
6482////////////////////////////////////////////////////////////////////////////////
6483
6484static void ggml_compute_backward(struct ggml_context * ctx, struct ggml_tensor * tensor, bool inplace) {
6485    struct ggml_tensor * src0 = tensor->src0;
6486    struct ggml_tensor * src1 = tensor->src1;
6487
6488    switch (tensor->op) {
6489        case GGML_OP_DUP:
6490            {
6491                if (src0->grad) {
6492                    src0->grad = ggml_add_impl(ctx, src0->grad, tensor->grad, inplace);
6493                }
6494            } break;
6495        case GGML_OP_ADD:
6496            {
6497                if (src0->grad) {
6498                    src0->grad = ggml_add_impl(ctx, src0->grad, tensor->grad, inplace);
6499                }
6500                if (src1->grad) {
6501                    src1->grad = ggml_add_impl(ctx, src1->grad, tensor->grad, inplace);
6502                }
6503            } break;
6504        case GGML_OP_SUB:
6505            {
6506                if (src0->grad) {
6507                    src0->grad = ggml_add_impl(ctx, src0->grad, tensor->grad, inplace);
6508                }
6509                if (src1->grad) {
6510                    src1->grad = ggml_sub_impl(ctx, src1->grad, tensor->grad, inplace);
6511                }
6512            } break;
6513        case GGML_OP_MUL:
6514            {
6515                if (src0->grad) {
6516                    src0->grad =
6517                        ggml_add_impl(ctx,
6518                                src0->grad,
6519                                ggml_mul(ctx, src1, tensor->grad),
6520                                inplace);
6521                }
6522                if (src1->grad) {
6523                    src1->grad =
6524                        ggml_add_impl(ctx,
6525                                src1->grad,
6526                                ggml_mul(ctx, src0, tensor->grad),
6527                                inplace);
6528                }
6529            } break;
6530        case GGML_OP_DIV:
6531            {
6532                if (src0->grad) {
6533                    src0->grad =
6534                        ggml_add_impl(ctx,
6535                                src0->grad,
6536                                ggml_div(ctx, tensor->grad, src1),
6537                                inplace);
6538                }
6539                if (src1->grad) {
6540                    src1->grad =
6541                        ggml_sub_impl(ctx,
6542                                src1->grad,
6543                                ggml_mul(ctx,
6544                                    tensor->grad,
6545                                    ggml_div(ctx, tensor, src1)),
6546                                inplace);
6547                }
6548            } break;
6549        case GGML_OP_SQR:
6550            {
6551                if (src0->grad) {
6552                    src0->grad =
6553                        ggml_add_impl(ctx,
6554                                src0->grad,
6555                                ggml_mul(ctx,
6556                                    ggml_mul(ctx, src0, tensor->grad),
6557                                    ggml_repeat(ctx, ggml_new_f32(ctx, 2.0f), src0)),
6558                                inplace);
6559                }
6560            } break;
6561        case GGML_OP_SQRT:
6562            {
6563                if (src0->grad) {
6564                    src0->grad =
6565                        ggml_add_impl(ctx,
6566                                src0->grad,
6567                                ggml_div(ctx,
6568                                    ggml_repeat(ctx, ggml_new_f32(ctx, 0.5f), tensor),
6569                                    tensor),
6570                                inplace);
6571                }
6572            } break;
6573        case GGML_OP_SUM:
6574            {
6575                if (src0->grad) {
6576                    src0->grad =
6577                        ggml_add_impl(ctx,
6578                                src0->grad,
6579                                ggml_repeat(ctx, tensor->grad, src0->grad),
6580                                inplace);
6581                }
6582            } break;
6583        case GGML_OP_MEAN:
6584            {
6585                assert(false); // TODO: implement
6586            } break;
6587        case GGML_OP_REPEAT:
6588            {
6589                if (src0->grad) {
6590                    src0->grad =
6591                        ggml_add_impl(ctx,
6592                                src0->grad,
6593                                ggml_sum(ctx, tensor->grad),
6594                                inplace);
6595                }
6596            } break;
6597        case GGML_OP_ABS:
6598            {
6599                if (src0->grad) {
6600                    src0->grad =
6601                        ggml_add_impl(ctx,
6602                                src0->grad,
6603                                ggml_mul(ctx,
6604                                    ggml_sgn(ctx, src0),
6605                                    tensor->grad),
6606                                inplace);
6607                }
6608            } break;
6609        case GGML_OP_SGN:
6610            {
6611                if (src0->grad) {
6612                    // noop
6613                }
6614            } break;
6615        case GGML_OP_NEG:
6616            {
6617                if (src0->grad) {
6618                    src0->grad = ggml_sub_impl(ctx, src0->grad, tensor->grad, inplace);
6619                }
6620            } break;
6621        case GGML_OP_STEP:
6622            {
6623                if (src0->grad) {
6624                    // noop
6625                }
6626            } break;
6627        case GGML_OP_RELU:
6628            {
6629                if (src0->grad) {
6630                    src0->grad = ggml_sub_impl(ctx,
6631                            src0->grad,
6632                            ggml_mul(ctx,
6633                                ggml_step(ctx, src0),
6634                                tensor->grad),
6635                            inplace);
6636                }
6637            } break;
6638        case GGML_OP_GELU:
6639            {
6640                assert(false); // TODO: not implemented
6641            } break;
6642        case GGML_OP_NORM:
6643            {
6644                assert(false); // TODO: not implemented
6645            } break;
6646        case GGML_OP_MUL_MAT:
6647            {
6648                if (src0->grad) {
6649                    // TODO: this requires outer product - ggml_out_prod(ctx, src1, tensor->grad);
6650                    assert(false);
6651                }
6652                if (src1->grad) {
6653                    src1->grad =
6654                        ggml_add_impl(ctx,
6655                                src1->grad,
6656                                // TODO: fix transpose, the node will break the graph connections
6657                                ggml_mul_mat(ctx, ggml_transpose(ctx, src0), tensor->grad),
6658                                inplace);
6659                }
6660            } break;
6661        case GGML_OP_SCALE:
6662            {
6663                GGML_ASSERT(false); // TODO: not implemented
6664            } break;
6665        case GGML_OP_CPY:
6666            {
6667                GGML_ASSERT(false); // TODO: not implemented
6668            } break;
6669        case GGML_OP_RESHAPE:
6670            {
6671                GGML_ASSERT(false); // TODO: not implemented
6672            } break;
6673        case GGML_OP_VIEW:
6674            {
6675                GGML_ASSERT(false); // not supported
6676            } break;
6677        case GGML_OP_PERMUTE:
6678            {
6679                GGML_ASSERT(false); // TODO: not implemented
6680            } break;
6681        case GGML_OP_TRANSPOSE:
6682            {
6683                GGML_ASSERT(false); // TODO: not implemented
6684            } break;
6685        case GGML_OP_GET_ROWS:
6686            {
6687                GGML_ASSERT(false); // TODO: not implemented
6688            } break;
6689        case GGML_OP_DIAG_MASK_INF:
6690            {
6691                GGML_ASSERT(false); // TODO: not implemented
6692            } break;
6693        case GGML_OP_SOFT_MAX:
6694            {
6695                GGML_ASSERT(false); // TODO: not implemented
6696            } break;
6697        case GGML_OP_ROPE:
6698            {
6699                GGML_ASSERT(false); // TODO: not implemented
6700            } break;
6701        case GGML_OP_CONV_1D_1S:
6702            {
6703                GGML_ASSERT(false); // TODO: not implemented
6704            } break;
6705        case GGML_OP_CONV_1D_2S:
6706            {
6707                GGML_ASSERT(false); // TODO: not implemented
6708            } break;
6709        case GGML_OP_FLASH_ATTN:
6710            {
6711                GGML_ASSERT(false); // not supported
6712            } break;
6713        case GGML_OP_FLASH_FF:
6714            {
6715                GGML_ASSERT(false); // not supported
6716            } break;
6717        case GGML_OP_NONE:
6718            {
6719                // nop
6720            } break;
6721        case GGML_OP_COUNT:
6722            {
6723                GGML_ASSERT(false);
6724            } break;
6725    }
6726}
6727
6728static void ggml_visit_parents(struct ggml_cgraph * cgraph, struct ggml_tensor * node) {
6729    if (node->grad == NULL) {
6730        // this usually happens when we generate intermediate nodes from constants in the backward pass
6731        // it can also happen during forward pass, if the user performs computations with constants
6732        if (node->op != GGML_OP_NONE) {
6733            //GGML_PRINT_DEBUG("%s: warning: node %p has no grad, but op %d\n", __func__, (void *) node, node->op);
6734        }
6735    }
6736
6737    // check if already visited
6738    for (int i = 0; i < cgraph->n_nodes; i++) {
6739        if (cgraph->nodes[i] == node) {
6740            return;
6741        }
6742    }
6743
6744    for (int i = 0; i < cgraph->n_leafs; i++) {
6745        if (cgraph->leafs[i] == node) {
6746            return;
6747        }
6748    }
6749
6750    if (node->src0) {
6751        ggml_visit_parents(cgraph, node->src0);
6752    }
6753
6754    if (node->src1) {
6755        ggml_visit_parents(cgraph, node->src1);
6756    }
6757
6758    for (int i = 0; i < GGML_MAX_OPT; ++i) {
6759        if (node->opt[i]) {
6760            ggml_visit_parents(cgraph, node->opt[i]);
6761        }
6762    }
6763
6764    if (node->op == GGML_OP_NONE && node->grad == NULL) {
6765        // reached a leaf node, not part of the gradient graph (e.g. a constant)
6766        assert(cgraph->n_leafs < GGML_MAX_NODES);
6767
6768        cgraph->leafs[cgraph->n_leafs] = node;
6769        cgraph->n_leafs++;
6770    } else {
6771        assert(cgraph->n_nodes < GGML_MAX_NODES);
6772
6773        cgraph->nodes[cgraph->n_nodes] = node;
6774        cgraph->grads[cgraph->n_nodes] = node->grad;
6775        cgraph->n_nodes++;
6776    }
6777}
6778
6779static void ggml_build_forward_impl(struct ggml_cgraph * cgraph, struct ggml_tensor * tensor, bool expand) {
6780    if (!expand) {
6781        cgraph->n_nodes = 0;
6782        cgraph->n_leafs = 0;
6783    }
6784
6785    const int n0 = cgraph->n_nodes;
6786    UNUSED(n0);
6787
6788    ggml_visit_parents(cgraph, tensor);
6789
6790    const int n_new = cgraph->n_nodes - n0;
6791    GGML_PRINT_DEBUG("%s: visited %d new nodes\n", __func__, n_new);
6792
6793    if (n_new > 0) {
6794        // the last added node should always be starting point
6795        assert(cgraph->nodes[cgraph->n_nodes - 1] == tensor);
6796    }
6797}
6798
6799void ggml_build_forward_expand(struct ggml_cgraph * cgraph, struct ggml_tensor * tensor) {
6800    ggml_build_forward_impl(cgraph, tensor, true);
6801}
6802
6803struct ggml_cgraph ggml_build_forward(struct ggml_tensor * tensor) {
6804    struct ggml_cgraph result = {
6805        /*.n_nodes      =*/ 0,
6806        /*.n_leafs      =*/ 0,
6807        /*.n_threads    =*/ 0,
6808        /*.work_size    =*/ 0,
6809        /*.work         =*/ NULL,
6810        /*.nodes        =*/ { NULL },
6811        /*.grads        =*/ { NULL },
6812        /*.leafs        =*/ { NULL },
6813        /*.perf_runs    =*/ 0,
6814        /*.perf_cycles  =*/ 0,
6815        /*.perf_time_us =*/ 0,
6816    };
6817
6818    ggml_build_forward_impl(&result, tensor, false);
6819
6820    return result;
6821}
6822
6823struct ggml_cgraph ggml_build_backward(struct ggml_context * ctx, struct ggml_cgraph * gf, bool keep) {
6824    struct ggml_cgraph result = *gf;
6825
6826    assert(gf->n_nodes > 0);
6827
6828    // if we are keeping the gradient graph, we have to detach the gradient nodes from the original graph
6829    if (keep) {
6830        for (int i = 0; i < gf->n_nodes; i++) {
6831            struct ggml_tensor * node = gf->nodes[i];
6832
6833            if (node->grad) {
6834                node->grad = ggml_dup_tensor(ctx, node);
6835                gf->grads[i] = node->grad;
6836            }
6837        }
6838    }
6839
6840    for (int i = gf->n_nodes - 1; i >= 0; i--) {
6841        struct ggml_tensor * node = gf->nodes[i];
6842
6843        // because we detached the grad nodes from the original graph, we can afford inplace operations
6844        if (node->grad) {
6845            ggml_compute_backward(ctx, node, keep);
6846        }
6847    }
6848
6849    for (int i = gf->n_nodes - 1; i >= 0; i--) {
6850        struct ggml_tensor * node = gf->nodes[i];
6851
6852        if (node->is_param) {
6853            GGML_PRINT_DEBUG("%s: found root node %p\n", __func__, (void *) node);
6854            ggml_build_forward_impl(&result, node->grad, true);
6855        }
6856    }
6857
6858    return result;
6859}
6860
6861//
6862// thread data
6863//
6864// synchronization is done via busy loops
6865// I tried using spin locks, but not sure how to use them correctly - the things I tried were slower than busy loops
6866//
6867
6868#ifdef __APPLE__
6869
6870//#include <os/lock.h>
6871
6872//typedef os_unfair_lock ggml_lock_t;
6873//
6874//#define ggml_lock_init(x)    UNUSED(x)
6875//#define ggml_lock_destroy(x) UNUSED(x)
6876//#define ggml_lock_lock       os_unfair_lock_lock
6877//#define ggml_lock_unlock     os_unfair_lock_unlock
6878//
6879//#define GGML_LOCK_INITIALIZER OS_UNFAIR_LOCK_INIT
6880
6881typedef int ggml_lock_t;
6882
6883#define ggml_lock_init(x)    UNUSED(x)
6884#define ggml_lock_destroy(x) UNUSED(x)
6885#define ggml_lock_lock(x)    UNUSED(x)
6886#define ggml_lock_unlock(x)  UNUSED(x)
6887
6888#define GGML_LOCK_INITIALIZER 0
6889
6890typedef pthread_t ggml_thread_t;
6891
6892#define ggml_thread_create pthread_create
6893#define ggml_thread_join   pthread_join
6894
6895#else
6896
6897//typedef pthread_spinlock_t ggml_lock_t;
6898
6899//#define ggml_lock_init(x) pthread_spin_init(x, PTHREAD_PROCESS_PRIVATE)
6900//#define ggml_lock_destroy pthread_spin_destroy
6901//#define ggml_lock_lock    pthread_spin_lock
6902//#define ggml_lock_unlock  pthread_spin_unlock
6903
6904typedef int ggml_lock_t;
6905
6906#define ggml_lock_init(x)    UNUSED(x)
6907#define ggml_lock_destroy(x) UNUSED(x)
6908#define ggml_lock_lock(x)    UNUSED(x)
6909#define ggml_lock_unlock(x)  UNUSED(x)
6910
6911#define GGML_LOCK_INITIALIZER 0
6912
6913typedef pthread_t ggml_thread_t;
6914
6915#define ggml_thread_create pthread_create
6916#define ggml_thread_join   pthread_join
6917
6918#endif
6919
6920struct ggml_compute_state_shared {
6921    ggml_lock_t spin;
6922
6923    int n_threads;
6924
6925    // synchronization primitives
6926    atomic_int  n_ready;
6927    atomic_bool has_work;
6928    atomic_bool stop; // stop all threads
6929};
6930
6931struct ggml_compute_state {
6932    ggml_thread_t thrd;
6933
6934    struct ggml_compute_params params;
6935    struct ggml_tensor * node;
6936
6937    struct ggml_compute_state_shared * shared;
6938};
6939
6940static thread_ret_t ggml_graph_compute_thread(void * data) {
6941    struct ggml_compute_state * state = (struct ggml_compute_state *) data;
6942
6943    const int n_threads = state->shared->n_threads;
6944
6945    while (true) {
6946        if (atomic_fetch_add(&state->shared->n_ready, 1) == n_threads - 1) {
6947            atomic_store(&state->shared->has_work, false);
6948        } else {
6949            while (atomic_load(&state->shared->has_work)) {
6950                if (atomic_load(&state->shared->stop)) {
6951                    return 0;
6952                }
6953                ggml_lock_lock  (&state->shared->spin);
6954                ggml_lock_unlock(&state->shared->spin);
6955            }
6956        }
6957
6958        atomic_fetch_sub(&state->shared->n_ready, 1);
6959
6960        // wait for work
6961        while (!atomic_load(&state->shared->has_work)) {
6962            if (atomic_load(&state->shared->stop)) {
6963                return 0;
6964            }
6965            ggml_lock_lock  (&state->shared->spin);
6966            ggml_lock_unlock(&state->shared->spin);
6967        }
6968
6969        // check if we should stop
6970        if (atomic_load(&state->shared->stop)) {
6971            break;
6972        }
6973
6974        if (state->node) {
6975            ggml_compute_forward(&state->params, state->node);
6976            state->node = NULL;
6977        } else {
6978            break;
6979        }
6980    }
6981
6982    return 0;
6983}
6984
6985void ggml_graph_compute(struct ggml_context * ctx, struct ggml_cgraph * cgraph) {
6986    if (cgraph->n_threads <= 0) {
6987        cgraph->n_threads = 8;
6988    }
6989
6990    const int n_threads = cgraph->n_threads;
6991
6992    struct ggml_compute_state_shared state_shared = {
6993        /*.spin      =*/ GGML_LOCK_INITIALIZER,
6994        /*.n_threads =*/ n_threads,
6995        /*.n_ready   =*/ 0,
6996        /*.has_work  =*/ false,
6997        /*.stop      =*/ false,
6998    };
6999    struct ggml_compute_state * workers = n_threads > 1 ? alloca(sizeof(struct ggml_compute_state)*(n_threads - 1)) : NULL;
7000
7001    // create thread pool
7002    if (n_threads > 1) {
7003        ggml_lock_init(&state_shared.spin);
7004
7005        atomic_store(&state_shared.has_work, true);
7006
7007        for (int j = 0; j < n_threads - 1; j++) {
7008            workers[j] = (struct ggml_compute_state) {
7009                .thrd   = 0,
7010                .params = {
7011                    .type  = GGML_TASK_COMPUTE,
7012                    .ith   = j + 1,
7013                    .nth   = n_threads,
7014                    .wsize = cgraph->work ? ggml_nbytes(cgraph->work) : 0,
7015                    .wdata = cgraph->work ? cgraph->work->data : NULL,
7016                },
7017                .node   = NULL,
7018                .shared = &state_shared,
7019            };
7020            int rc = ggml_thread_create(&workers[j].thrd, NULL, ggml_graph_compute_thread, &workers[j]);
7021            assert(rc == 0);
7022            UNUSED(rc);
7023        }
7024    }
7025
7026    // initialize tasks + work buffer
7027    {
7028        size_t work_size = 0;
7029
7030        // thread scheduling for the different operations
7031        for (int i = 0; i < cgraph->n_nodes; i++) {
7032            struct ggml_tensor * node = cgraph->nodes[i];
7033
7034            switch (node->op) {
7035                case GGML_OP_DUP:
7036                    {
7037                        node->n_tasks = 1;
7038                    } break;
7039                case GGML_OP_ADD:
7040                    {
7041                        node->n_tasks = n_threads;
7042                    } break;
7043                case GGML_OP_SUB:
7044                case GGML_OP_MUL:
7045                case GGML_OP_DIV:
7046                case GGML_OP_SQR:
7047                case GGML_OP_SQRT:
7048                case GGML_OP_SUM:
7049                case GGML_OP_MEAN:
7050                case GGML_OP_REPEAT:
7051                case GGML_OP_ABS:
7052                case GGML_OP_SGN:
7053                case GGML_OP_NEG:
7054                case GGML_OP_STEP:
7055                case GGML_OP_RELU:
7056                    {
7057                        node->n_tasks = 1;
7058                    } break;
7059                case GGML_OP_GELU:
7060                    {
7061                        node->n_tasks = n_threads;
7062                    } break;
7063                case GGML_OP_NORM:
7064                    {
7065                        node->n_tasks = n_threads;
7066                    } break;
7067                case GGML_OP_MUL_MAT:
7068                    {
7069                        // TODO: use different scheduling for different matrix sizes
7070                        node->n_tasks = n_threads;
7071
7072                        size_t cur = 0;
7073
7074                        // TODO: better way to determine if the matrix is transposed
7075                        if (node->src0->nb[1] < node->src0->nb[0]) {
7076                            cur = ggml_nbytes(node)*node->n_tasks; // TODO: this can become (n_tasks-1)
7077                        } else {
7078                            if (node->src0->type == GGML_TYPE_F16 &&
7079                                node->src1->type == GGML_TYPE_F32) {
7080#if defined(GGML_USE_ACCELERATE) || defined(GGML_USE_OPENBLAS)
7081                                if (ggml_compute_forward_mul_mat_use_blas(node->src0, node->src1, node)) {
7082                                    cur = sizeof(float)*(node->src0->ne[0]*node->src0->ne[1]);
7083                                } else {
7084                                    cur = sizeof(ggml_fp16_t)*ggml_nelements(node->src1);
7085                                }
7086#else
7087                                cur = sizeof(ggml_fp16_t)*ggml_nelements(node->src1);
7088#endif
7089                            } else if (node->src0->type == GGML_TYPE_F32 &&
7090                                       node->src1->type == GGML_TYPE_F32) {
7091                                cur = 0;
7092                            } else {
7093                                GGML_ASSERT(false);
7094                            }
7095                        }
7096
7097                        work_size = MAX(work_size, cur);
7098                    } break;
7099                case GGML_OP_SCALE:
7100                    {
7101                        node->n_tasks = n_threads;
7102                    } break;
7103                case GGML_OP_CPY:
7104                case GGML_OP_RESHAPE:
7105                case GGML_OP_VIEW:
7106                case GGML_OP_PERMUTE:
7107                case GGML_OP_TRANSPOSE:
7108                case GGML_OP_GET_ROWS:
7109                case GGML_OP_DIAG_MASK_INF:
7110                    {
7111                        node->n_tasks = 1;
7112                    } break;
7113                case GGML_OP_SOFT_MAX:
7114                    {
7115                        node->n_tasks = n_threads;
7116                    } break;
7117                case GGML_OP_ROPE:
7118                    {
7119                        node->n_tasks = 1;
7120                    } break;
7121                case GGML_OP_CONV_1D_1S:
7122                case GGML_OP_CONV_1D_2S:
7123                    {
7124                        node->n_tasks = n_threads;
7125
7126                        GGML_ASSERT(node->src0->ne[3] == 1);
7127                        GGML_ASSERT(node->src1->ne[2] == 1);
7128                        GGML_ASSERT(node->src1->ne[3] == 1);
7129
7130                        size_t cur = 0;
7131                        const int nk = node->src0->ne[0];
7132
7133                        if (node->src0->type == GGML_TYPE_F16 &&
7134                            node->src1->type == GGML_TYPE_F32) {
7135                            cur = sizeof(ggml_fp16_t)*(
7136                                    nk*ggml_up32(node->src0->ne[1])*node->src0->ne[2] +
7137                                    ( 2*(nk/2) + node->src1->ne[0])*node->src1->ne[1]
7138                                    );
7139                        } else if (node->src0->type == GGML_TYPE_F32 &&
7140                                   node->src1->type == GGML_TYPE_F32) {
7141                            cur = sizeof(float)*(
7142                                    nk*ggml_up32(node->src0->ne[1])*node->src0->ne[2] +
7143                                    ( 2*(nk/2) + node->src1->ne[0])*node->src1->ne[1]
7144                                    );
7145                        } else {
7146                            GGML_ASSERT(false);
7147                        }
7148
7149                        work_size = MAX(work_size, cur);
7150                    } break;
7151                case GGML_OP_FLASH_ATTN:
7152                    {
7153                        node->n_tasks = n_threads;
7154
7155                        size_t cur = 0;
7156
7157                        if (node->src1->type == GGML_TYPE_F32) {
7158                            cur  = sizeof(float)*node->src1->ne[1]*node->n_tasks; // TODO: this can become (n_tasks-1)
7159                            cur += sizeof(float)*node->src1->ne[1]*node->n_tasks; // this is overestimated by x2
7160                        }
7161
7162                        if (node->src1->type == GGML_TYPE_F16) {
7163                            cur  = sizeof(float)*node->src1->ne[1]*node->n_tasks; // TODO: this can become (n_tasks-1)
7164                            cur += sizeof(float)*node->src1->ne[1]*node->n_tasks; // this is overestimated by x2
7165                        }
7166
7167                        work_size = MAX(work_size, cur);
7168                    } break;
7169                case GGML_OP_FLASH_FF:
7170                    {
7171                        node->n_tasks = n_threads;
7172
7173                        size_t cur = 0;
7174
7175                        if (node->src1->type == GGML_TYPE_F32) {
7176                            cur  = sizeof(float)*node->src1->ne[1]*node->n_tasks; // TODO: this can become (n_tasks-1)
7177                            cur += sizeof(float)*node->src1->ne[1]*node->n_tasks; // this is overestimated by x2
7178                        }
7179
7180                        if (node->src1->type == GGML_TYPE_F16) {
7181                            cur  = sizeof(float)*node->src1->ne[1]*node->n_tasks; // TODO: this can become (n_tasks-1)
7182                            cur += sizeof(float)*node->src1->ne[1]*node->n_tasks; // this is overestimated by x2
7183                        }
7184
7185                        work_size = MAX(work_size, cur);
7186                    } break;
7187                case GGML_OP_NONE:
7188                    {
7189                        node->n_tasks = 1;
7190                    } break;
7191                case GGML_OP_COUNT:
7192                    {
7193                        assert(false);
7194                    } break;
7195            }
7196        }
7197
7198        if (cgraph->work != NULL && work_size > cgraph->work_size) {
7199            assert(false); // TODO: better handling
7200        }
7201
7202        if (work_size > 0 && cgraph->work == NULL) {
7203            cgraph->work_size = work_size + CACHE_LINE_SIZE*(n_threads - 1);
7204
7205            GGML_PRINT_DEBUG("%s: allocating work buffer for graph (%zu bytes)\n", __func__, cgraph->work_size);
7206            cgraph->work = ggml_new_tensor_1d(ctx, GGML_TYPE_I8, cgraph->work_size);
7207        }
7208    }
7209
7210    const int64_t perf_start_cycles  = ggml_perf_cycles();
7211    const int64_t perf_start_time_us = ggml_perf_time_us();
7212
7213    for (int i = 0; i < cgraph->n_nodes; i++) {
7214        GGML_PRINT_DEBUG_5("%s: %d/%d\n", __func__, i, cgraph->n_nodes);
7215
7216        struct ggml_tensor * node = cgraph->nodes[i];
7217
7218        // TODO: this could be used to avoid unnecessary computations, but it needs to be improved
7219        //if (node->grad == NULL && node->perf_runs > 0) {
7220        //    continue;
7221        //}
7222
7223        const int64_t perf_node_start_cycles  = ggml_perf_cycles();
7224        const int64_t perf_node_start_time_us = ggml_perf_time_us();
7225
7226        // INIT
7227        struct ggml_compute_params params = {
7228            /*.type  =*/ GGML_TASK_INIT,
7229            /*.ith   =*/ 0,
7230            /*.nth   =*/ node->n_tasks,
7231            /*.wsize =*/ cgraph->work ? ggml_nbytes(cgraph->work) : 0,
7232            /*.wdata =*/ cgraph->work ? cgraph->work->data : NULL,
7233        };
7234
7235        ggml_compute_forward(&params, node);
7236
7237        // COMPUTE
7238        if (node->n_tasks > 1) {
7239            if (atomic_fetch_add(&state_shared.n_ready, 1) == n_threads - 1) {
7240                atomic_store(&state_shared.has_work, false);
7241            }
7242
7243            while (atomic_load(&state_shared.has_work)) {
7244                ggml_lock_lock  (&state_shared.spin);
7245                ggml_lock_unlock(&state_shared.spin);
7246            }
7247
7248            // launch thread pool
7249            for (int j = 0; j < n_threads - 1; j++) {
7250                workers[j].params = (struct ggml_compute_params) {
7251                    .type  = GGML_TASK_COMPUTE,
7252                    .ith   = j + 1,
7253                    .nth   = n_threads,
7254                    .wsize = cgraph->work ? ggml_nbytes(cgraph->work) : 0,
7255                    .wdata = cgraph->work ? cgraph->work->data : NULL,
7256                };
7257                workers[j].node = node;
7258            }
7259
7260            atomic_fetch_sub(&state_shared.n_ready, 1);
7261
7262            while (atomic_load(&state_shared.n_ready) > 0) {
7263                ggml_lock_lock  (&state_shared.spin);
7264                ggml_lock_unlock(&state_shared.spin);
7265            }
7266
7267            atomic_store(&state_shared.has_work, true);
7268        }
7269
7270        params.type = GGML_TASK_COMPUTE;
7271        ggml_compute_forward(&params, node);
7272
7273        // wait for thread pool
7274        if (node->n_tasks > 1) {
7275            if (atomic_fetch_add(&state_shared.n_ready, 1) == n_threads - 1) {
7276                atomic_store(&state_shared.has_work, false);
7277            }
7278
7279            while (atomic_load(&state_shared.has_work)) {
7280                ggml_lock_lock  (&state_shared.spin);
7281                ggml_lock_unlock(&state_shared.spin);
7282            }
7283
7284            atomic_fetch_sub(&state_shared.n_ready, 1);
7285
7286            while (atomic_load(&state_shared.n_ready) != 0) {
7287                ggml_lock_lock  (&state_shared.spin);
7288                ggml_lock_unlock(&state_shared.spin);
7289            }
7290        }
7291
7292        // FINALIZE
7293        if (node->n_tasks > 1) {
7294            if (atomic_fetch_add(&state_shared.n_ready, 1) == n_threads - 1) {
7295                atomic_store(&state_shared.has_work, false);
7296            }
7297
7298            while (atomic_load(&state_shared.has_work)) {
7299                ggml_lock_lock  (&state_shared.spin);
7300                ggml_lock_unlock(&state_shared.spin);
7301            }
7302
7303            // launch thread pool
7304            for (int j = 0; j < n_threads - 1; j++) {
7305                workers[j].params = (struct ggml_compute_params) {
7306                    .type  = GGML_TASK_FINALIZE,
7307                    .ith   = j + 1,
7308                    .nth   = n_threads,
7309                    .wsize = cgraph->work ? ggml_nbytes(cgraph->work) : 0,
7310                    .wdata = cgraph->work ? cgraph->work->data : NULL,
7311                };
7312                workers[j].node = node;
7313            }
7314
7315            atomic_fetch_sub(&state_shared.n_ready, 1);
7316
7317            while (atomic_load(&state_shared.n_ready) > 0) {
7318                ggml_lock_lock  (&state_shared.spin);
7319                ggml_lock_unlock(&state_shared.spin);
7320            }
7321
7322            atomic_store(&state_shared.has_work, true);
7323        }
7324
7325        params.type = GGML_TASK_FINALIZE;
7326        ggml_compute_forward(&params, node);
7327
7328        // wait for thread pool
7329        if (node->n_tasks > 1) {
7330            if (atomic_fetch_add(&state_shared.n_ready, 1) == n_threads - 1) {
7331                atomic_store(&state_shared.has_work, false);
7332            }
7333
7334            while (atomic_load(&state_shared.has_work)) {
7335                ggml_lock_lock  (&state_shared.spin);
7336                ggml_lock_unlock(&state_shared.spin);
7337            }
7338
7339            atomic_fetch_sub(&state_shared.n_ready, 1);
7340
7341            while (atomic_load(&state_shared.n_ready) != 0) {
7342                ggml_lock_lock  (&state_shared.spin);
7343                ggml_lock_unlock(&state_shared.spin);
7344            }
7345        }
7346
7347        // performance stats (node)
7348        {
7349            int64_t perf_cycles_cur  = ggml_perf_cycles()  - perf_node_start_cycles;
7350            int64_t perf_time_us_cur = ggml_perf_time_us() - perf_node_start_time_us;
7351
7352            node->perf_runs++;
7353            node->perf_cycles  += perf_cycles_cur;
7354            node->perf_time_us += perf_time_us_cur;
7355        }
7356    }
7357
7358    // join thread pool
7359    if (n_threads > 1) {
7360        atomic_store(&state_shared.stop, true);
7361        atomic_store(&state_shared.has_work, true);
7362
7363        for (int j = 0; j < n_threads - 1; j++) {
7364            int rc = ggml_thread_join(workers[j].thrd, NULL);
7365            assert(rc == 0);
7366            UNUSED(rc);
7367        }
7368
7369        ggml_lock_destroy(&state_shared.spin);
7370    }
7371
7372    // performance stats (graph)
7373    {
7374        int64_t perf_cycles_cur  = ggml_perf_cycles()  - perf_start_cycles;
7375        int64_t perf_time_us_cur = ggml_perf_time_us() - perf_start_time_us;
7376
7377        cgraph->perf_runs++;
7378        cgraph->perf_cycles  += perf_cycles_cur;
7379        cgraph->perf_time_us += perf_time_us_cur;
7380
7381        GGML_PRINT_DEBUG("%s: perf (%d) - cpu = %.3f / %.3f ms, wall = %.3f / %.3f ms\n",
7382                __func__, cgraph->perf_runs,
7383                (double) perf_cycles_cur      / (double) ggml_cycles_per_ms(),
7384                (double) cgraph->perf_cycles  / (double) ggml_cycles_per_ms() / (double) cgraph->perf_runs,
7385                (double) perf_time_us_cur     / 1000.0,
7386                (double) cgraph->perf_time_us / 1000.0 / cgraph->perf_runs);
7387    }
7388}
7389
7390void ggml_graph_reset(struct ggml_cgraph * cgraph) {
7391    for (int i = 0; i < cgraph->n_nodes; i++) {
7392        struct ggml_tensor * grad = cgraph->grads[i];
7393
7394        if (grad) {
7395            ggml_set_zero(grad);
7396        }
7397    }
7398}
7399
7400void ggml_graph_print(const struct ggml_cgraph * cgraph) {
7401    int64_t perf_total_per_op_us[GGML_OP_COUNT] = {0};
7402
7403    GGML_PRINT("=== GRAPH ===\n");
7404
7405    GGML_PRINT_DEBUG("n_threads       = %d\n",       cgraph->n_threads);
7406    GGML_PRINT_DEBUG("total work size = %zu bytes\n",cgraph->work_size);
7407
7408    GGML_PRINT("n_nodes = %d\n", cgraph->n_nodes);
7409    for (int i = 0; i < cgraph->n_nodes; i++) {
7410        struct ggml_tensor * node = cgraph->nodes[i];
7411
7412        perf_total_per_op_us[node->op] += node->perf_time_us;
7413
7414        GGML_PRINT(" - %3d: [ %6d, %6d, %6d] %16s %s (%3d) cpu = %7.3f / %7.3f ms, wall = %7.3f / %7.3f ms\n",
7415                i,
7416                node->ne[0], node->ne[1], node->ne[2],
7417                GGML_OP_LABEL[node->op], node->is_param ? "x" : node->grad ? "g" : " ", node->perf_runs,
7418                (double) node->perf_cycles  / (double) ggml_cycles_per_ms(),
7419                (double) node->perf_cycles  / (double) ggml_cycles_per_ms() / (double) node->perf_runs,
7420                (double) node->perf_time_us / 1000.0,
7421                (double) node->perf_time_us / 1000.0 / node->perf_runs);
7422    }
7423
7424    GGML_PRINT("n_leafs = %d\n", cgraph->n_leafs);
7425    for (int i = 0; i < cgraph->n_leafs; i++) {
7426        struct ggml_tensor * node = cgraph->leafs[i];
7427
7428        GGML_PRINT(" - %3d: [ %6d, %6d] %8s\n",
7429                i,
7430                node->ne[0], node->ne[1],
7431                GGML_OP_LABEL[node->op]);
7432    }
7433
7434    for (int i = 0; i < GGML_OP_COUNT; i++) {
7435        GGML_PRINT("perf_total_per_op_us[%16s] = %7.3f ms\n", GGML_OP_LABEL[i], (double) perf_total_per_op_us[i] / 1000.0);
7436    }
7437
7438    GGML_PRINT("========================================\n");
7439}
7440
7441// check if node is part of the graph
7442static bool ggml_graph_find(const struct ggml_cgraph * cgraph, const struct ggml_tensor * node) {
7443    if (cgraph == NULL) {
7444        return true;
7445    }
7446
7447    for (int i = 0; i < cgraph->n_nodes; i++) {
7448        if (cgraph->nodes[i] == node) {
7449            return true;
7450        }
7451    }
7452
7453    return false;
7454}
7455
7456static struct ggml_tensor * ggml_graph_get_parent(const struct ggml_cgraph * cgraph, const struct ggml_tensor * node) {
7457    for (int i = 0; i < cgraph->n_nodes; i++) {
7458        struct ggml_tensor * parent = cgraph->nodes[i];
7459
7460        if (parent->grad == node) {
7461            return parent;
7462        }
7463    }
7464
7465    return NULL;
7466}
7467
7468void ggml_graph_dump_dot(const struct ggml_cgraph * gb, const struct ggml_cgraph * gf, const char * filename) {
7469    char color[16];
7470
7471    FILE * fp = fopen(filename, "w");
7472    assert(fp);
7473
7474    fprintf(fp, "digraph G {\n");
7475    fprintf(fp, "  newrank = true;\n");
7476    fprintf(fp, "  rankdir = LR;\n");
7477
7478    for (int i = 0; i < gb->n_nodes; i++) {
7479        struct ggml_tensor * node = gb->nodes[i];
7480
7481        if (ggml_graph_get_parent(gb, node) != NULL) {
7482            continue;
7483        }
7484
7485        if (node->is_param) {
7486            snprintf(color, sizeof(color), "yellow");
7487        } else if (node->grad) {
7488            if (ggml_graph_find(gf, node)) {
7489                snprintf(color, sizeof(color), "green");
7490            } else {
7491                snprintf(color, sizeof(color), "lightblue");
7492            }
7493        } else {
7494            snprintf(color, sizeof(color), "white");
7495        }
7496
7497        fprintf(fp, "  \"%p\" [ \
7498style = filled; fillcolor = %s; shape = record; \
7499label=\"%d [%d, %d] | <x>%s",
7500                (void *) node, color,
7501                i, node->ne[0], node->ne[1],
7502                GGML_OP_SYMBOL[node->op]);
7503
7504        if (node->grad) {
7505            fprintf(fp, " | <g>%s\"; ]\n", GGML_OP_SYMBOL[node->grad->op]);
7506        } else {
7507            fprintf(fp, "\"; ]\n");
7508        }
7509    }
7510
7511    for (int i = 0; i < gb->n_leafs; i++) {
7512        struct ggml_tensor * node = gb->leafs[i];
7513
7514        snprintf(color, sizeof(color), "pink");
7515
7516        if (ggml_nelements(node) == 1) {
7517            fprintf(fp, "  \"%p\" [ \
7518style = filled; fillcolor = %s; shape = record; \
7519label=\"<x>%.1e\"; ]\n",
7520                    (void *) node, color, ggml_get_f32_1d(node, 0));
7521        } else {
7522            fprintf(fp, "  \"%p\" [ \
7523style = filled; fillcolor = %s; shape = record; \
7524label=\"<x>CONST %d [%d, %d]\"; ]\n",
7525                    (void *) node, color,
7526                    i, node->ne[0], node->ne[1]);
7527        }
7528    }
7529
7530    for (int i = 0; i < gb->n_nodes; i++) {
7531        struct ggml_tensor * node = gb->nodes[i];
7532
7533        struct ggml_tensor * parent = ggml_graph_get_parent(gb, node);
7534
7535        if (node->src0) {
7536            struct ggml_tensor * parent0 = ggml_graph_get_parent(gb, node->src0);
7537
7538            fprintf(fp, "  \"%p\":%s -> \"%p\":%s [ arrowhead = %s; style = %s; label = \"x\"; ]\n",
7539                    parent0 ? (void *) parent0 : (void *) node->src0,
7540                    parent0 ? "g" : "x",
7541                    parent ? (void *) parent : (void *) node,
7542                    parent ? "g" : "x",
7543                    parent ? "empty" : "vee",
7544                    parent ? "dashed" : "solid");
7545        }
7546
7547        if (node->src1) {
7548            struct ggml_tensor * parent1 = ggml_graph_get_parent(gb, node->src1);
7549
7550            fprintf(fp, "  \"%p\":%s -> \"%p\":%s [ arrowhead = %s; style = %s; label = \"y\"; ]\n",
7551                    parent1 ? (void *) parent1 : (void *) node->src1,
7552                    parent1 ? "g" : "x",
7553                    parent ? (void *) parent : (void *) node,
7554                    parent ? "g" : "x",
7555                    parent ? "empty" : "vee",
7556                    parent ? "dashed" : "solid");
7557        }
7558    }
7559
7560    for (int i = 0; i < gb->n_leafs; i++) {
7561        struct ggml_tensor * node = gb->leafs[i];
7562
7563        if (node->src0) {
7564            fprintf(fp, "  \"%p\":%s -> \"%p\":%s [ label = \"x\"; ]\n",
7565                    (void *) node->src0, "x",
7566                    (void *) node, "x");
7567        }
7568
7569        if (node->src1) {
7570            fprintf(fp, "  \"%p\":%s -> \"%p\":%s [ label = \"y\"; ]\n",
7571                    (void *) node->src1, "x",
7572                    (void *) node, "x");
7573        }
7574    }
7575
7576    fprintf(fp, "}\n");
7577
7578    fclose(fp);
7579
7580    GGML_PRINT("%s: dot -Tpng %s -o %s.png && open %s.png\n", __func__, filename, filename, filename);
7581}
7582
7583////////////////////////////////////////////////////////////////////////////////
7584
7585static void ggml_opt_set_params(int np, struct ggml_tensor * const ps[], const float * x) {
7586    int i = 0;
7587    for (int p = 0; p < np; ++p) {
7588        const int ne = ggml_nelements(ps[p]) ;
7589        // TODO: add function to set tensor from array
7590        for (int j = 0; j < ne; ++j) {
7591            ggml_set_f32_1d(ps[p], j, x[i++]);
7592        }
7593    }
7594}
7595
7596static void ggml_opt_get_params(int np, struct ggml_tensor * const ps[], float * x) {
7597    int i = 0;
7598    for (int p = 0; p < np; ++p) {
7599        const int ne = ggml_nelements(ps[p]) ;
7600        // TODO: add function to get all elements at once
7601        for (int j = 0; j < ne; ++j) {
7602            x[i++] = ggml_get_f32_1d(ps[p], j);
7603        }
7604    }
7605}
7606
7607static void ggml_opt_get_grad(int np, struct ggml_tensor * const ps[], float * g) {
7608    int i = 0;
7609    for (int p = 0; p < np; ++p) {
7610        const int ne = ggml_nelements(ps[p]) ;
7611        // TODO: add function to get all elements at once
7612        for (int j = 0; j < ne; ++j) {
7613            g[i++] = ggml_get_f32_1d(ps[p]->grad, j);
7614        }
7615    }
7616}
7617
7618//
7619// ADAM
7620//
7621//   ref: https://arxiv.org/pdf/1412.6980.pdf
7622//
7623
7624static enum ggml_opt_result ggml_opt_adam(
7625        struct ggml_context * ctx,
7626        struct ggml_opt_params params,
7627        struct ggml_tensor * f,
7628        struct ggml_cgraph * gf,
7629        struct ggml_cgraph * gb) {
7630    assert(ggml_is_scalar(f));
7631
7632    gf->n_threads = params.n_threads;
7633    gb->n_threads = params.n_threads;
7634
7635    // these will store the parameters we want to optimize
7636    struct ggml_tensor * ps[GGML_MAX_PARAMS];
7637
7638    int np = 0;
7639    int nx = 0;
7640    for (int i = 0; i < gf->n_nodes; ++i) {
7641        if (gf->nodes[i]->is_param) {
7642            GGML_PRINT_DEBUG("found param %d: grad->op = %d\n", np, gf->nodes[i]->grad->op);
7643
7644            assert(np < GGML_MAX_PARAMS);
7645
7646            ps[np++] = gf->nodes[i];
7647            nx += ggml_nelements(gf->nodes[i]);
7648        }
7649    }
7650
7651    // constants
7652    const float alpha = params.adam.alpha;
7653    const float beta1 = params.adam.beta1;
7654    const float beta2 = params.adam.beta2;
7655    const float eps   = params.adam.eps;
7656
7657    float * x  = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data; // view of the parameters
7658    float * g1 = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data; // gradient
7659    float * g2 = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data; // gradient squared
7660    float * m  = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data; // first moment
7661    float * v  = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data; // second moment
7662    float * mh = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data; // first moment hat
7663    float * vh = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data; // second moment hat
7664
7665    float * pf = params.past > 0 ? ggml_new_tensor_1d(ctx, GGML_TYPE_F32, params.past)->data : NULL; // past function values
7666
7667    // initialize
7668    ggml_vec_set_f32(nx, m, 0.0f);
7669    ggml_vec_set_f32(nx, v, 0.0f);
7670
7671    // update view
7672    ggml_opt_get_params(np, ps, x);
7673
7674    // compute the function value
7675    ggml_graph_reset  (gf);
7676    ggml_set_f32      (f->grad, 1.0f);
7677    ggml_graph_compute(ctx, gb);
7678
7679    float fx_prev = ggml_get_f32_1d(f, 0);
7680    if (pf) {
7681        pf[0] = fx_prev;
7682    }
7683
7684    int n_no_improvement = 0;
7685    float fx_best = fx_prev;
7686
7687    // run the optimizer
7688    for (int t = 0; t < params.adam.n_iter; ++t) {
7689        GGML_PRINT_DEBUG  ("=== iter %d ===\n", t);
7690
7691        GGML_PRINT_DEBUG  ("f      = %10.6f\n", ggml_get_f32_1d(f, 0));
7692        GGML_PRINT_DEBUG_5("df/dx0 = %10.6f\n", ggml_get_f32_1d(ps[0]->grad, 0));
7693        GGML_PRINT_DEBUG_5("df/dx1 = %10.6f\n", ggml_get_f32_1d(ps[1]->grad, 0));
7694
7695        for (int i = 0; i < np; ++i) {
7696            GGML_PRINT_DEBUG("param %d: %10.6f, g = %10.6f\n", i,
7697                    ggml_get_f32_1d(ps[i], 0), ggml_get_f32_1d(ps[i]->grad, 0));
7698        }
7699
7700        const int64_t t_start_wall = ggml_time_us();
7701        const int64_t t_start_cpu = ggml_cycles();
7702        UNUSED(t_start_wall);
7703        UNUSED(t_start_cpu);
7704
7705        {
7706            // update the gradient
7707            ggml_opt_get_grad(np, ps, g1);
7708
7709            // m_t = beta1*m_t-1 + (1 - beta1)*g_t
7710            ggml_vec_scale_f32(nx, m, beta1);
7711            ggml_vec_mad_f32  (nx, m, g1, 1.0f - beta1);
7712
7713            // g2 = g1^2
7714            ggml_vec_sqr_f32  (nx, g2, g1);
7715
7716            // v_t = beta2*v_t-1 + (1 - beta2)*g_t^2
7717            ggml_vec_scale_f32(nx, v, beta2);
7718            ggml_vec_mad_f32  (nx, v, g2, 1.0f - beta2);
7719
7720            // m^hat = m_t / (1 - beta1^t)
7721            // v^hat = v_t / (1 - beta2^t)
7722            // x_t = x_t-1 - alpha*m^hat/(sqrt(v^hat) + eps)
7723            ggml_vec_cpy_f32  (nx, mh, m);
7724            ggml_vec_cpy_f32  (nx, vh, v);
7725
7726            ggml_vec_scale_f32(nx, mh, alpha/(1.0f - powf(beta1, t + 1)));
7727            ggml_vec_scale_f32(nx, vh,  1.0f/(1.0f - powf(beta2, t + 1)));
7728
7729            ggml_vec_sqrt_f32 (nx, vh, vh);
7730            ggml_vec_acc1_f32 (nx, vh, eps);
7731
7732            ggml_vec_div_f32  (nx, mh, mh, vh);
7733            ggml_vec_sub_f32  (nx, x,  x,  mh);
7734
7735            // update the parameters
7736            ggml_opt_set_params(np, ps, x);
7737        }
7738
7739        ggml_graph_reset  (gf);
7740        ggml_set_f32      (f->grad, 1.0f);
7741        ggml_graph_compute(ctx, gb);
7742
7743        const float fx = ggml_get_f32_1d(f, 0);
7744
7745        // check convergence
7746        if (fabsf(fx - fx_prev)/fx < params.adam.eps_f) {
7747            GGML_PRINT_DEBUG("converged\n");
7748
7749            return GGML_OPT_OK;
7750        }
7751
7752        // delta-based convergence test
7753        if (pf != NULL) {
7754            // need at least params.past iterations to start checking for convergence
7755            if (params.past <= t) {
7756                const float rate = (pf[t%params.past] - fx)/fx;
7757
7758                if (fabs(rate) < params.delta) {
7759                    return GGML_OPT_OK;
7760                }
7761            }
7762
7763            pf[t%params.past] = fx;
7764        }
7765
7766        // check for improvement
7767        if (params.max_no_improvement > 0) {
7768            if (fx_best > fx) {
7769                fx_best = fx;
7770                n_no_improvement = 0;
7771            } else {
7772                ++n_no_improvement;
7773
7774                if (n_no_improvement >= params.max_no_improvement) {
7775                    return GGML_OPT_OK;
7776                }
7777            }
7778        }
7779
7780        fx_prev = fx;
7781
7782        {
7783            const int64_t t_end_cpu = ggml_cycles();
7784            GGML_PRINT_DEBUG("time iter:      %5.3f s\n", ((float)(t_end_cpu - t_start_cpu))/CLOCKS_PER_SEC);
7785            UNUSED(t_end_cpu);
7786
7787            const int64_t t_end_wall = ggml_time_us();
7788            GGML_PRINT_DEBUG("wall time iter: %5.3f s\n", (t_end_wall - t_start_wall)/1e6);
7789            UNUSED(t_end_wall);
7790        }
7791    }
7792
7793    return GGML_OPT_DID_NOT_CONVERGE;
7794}
7795
7796//
7797// L-BFGS
7798//
7799// the L-BFGS implementation below is based on the following implementation:
7800//
7801//   https://github.com/chokkan/liblbfgs
7802//
7803
7804struct ggml_lbfgs_iteration_data {
7805    float alpha;
7806    float ys;
7807    float * s;
7808    float * y;
7809};
7810
7811static enum ggml_opt_result linesearch_backtracking(
7812        struct ggml_context * ctx,
7813        const struct ggml_opt_params * params,
7814        int nx,
7815        float * x,
7816        float * fx,
7817        float * g,
7818        float * d,
7819        float * step,
7820        const float * xp,
7821        struct ggml_tensor * f,
7822        struct ggml_cgraph * gf,
7823        struct ggml_cgraph * gb,
7824        const int np,
7825        struct ggml_tensor * ps[]) {
7826    int count = 0;
7827
7828    float width  = 0.0f;
7829    float dg     = 0.0f;
7830    float finit  = 0.0f;
7831    float dginit = 0.0f;
7832    float dgtest = 0.0f;
7833
7834    const float dec = 0.5f;
7835    const float inc = 2.1f;
7836
7837    if (*step <= 0.) {
7838        return GGML_LINESEARCH_INVALID_PARAMETERS;
7839    }
7840
7841    // compute the initial gradient in the search direction
7842    ggml_vec_dot_f32(nx, &dginit, g, d);
7843
7844    // make sure that d points to a descent direction
7845    if (0 < dginit) {
7846        return GGML_LINESEARCH_FAIL;
7847    }
7848
7849    // initialize local variables
7850    finit = *fx;
7851    dgtest = params->lbfgs.ftol*dginit;
7852
7853    while (true) {
7854        ggml_vec_cpy_f32(nx, x, xp);
7855        ggml_vec_mad_f32(nx, x, d, *step);
7856
7857        // evaluate the function and gradient values
7858        {
7859            ggml_opt_set_params(np, ps, x);
7860
7861            ggml_graph_reset  (gf);
7862            ggml_set_f32      (f->grad, 1.0f);
7863            ggml_graph_compute(ctx, gb);
7864
7865            ggml_opt_get_grad(np, ps, g);
7866
7867            *fx = ggml_get_f32_1d(f, 0);
7868        }
7869
7870        ++count;
7871
7872        if (*fx > finit + (*step)*dgtest) {
7873            width = dec;
7874        } else {
7875            // Armijo condition is satisfied
7876            if (params->lbfgs.linesearch == GGML_LINESEARCH_BACKTRACKING_ARMIJO) {
7877                return count;
7878            }
7879
7880            ggml_vec_dot_f32(nx, &dg, g, d);
7881
7882            // check the Wolfe condition
7883            if (dg < params->lbfgs.wolfe * dginit) {
7884                width = inc;
7885            } else {
7886                if(params->lbfgs.linesearch == GGML_LINESEARCH_BACKTRACKING_WOLFE) {
7887                    // regular Wolfe conditions
7888                    return count;
7889                }
7890
7891                if(dg > -params->lbfgs.wolfe*dginit) {
7892                    width = dec;
7893                } else {
7894                    // strong Wolfe condition (GGML_LINESEARCH_BACKTRACKING_STRONG_WOLFE)
7895                    return count;
7896                }
7897                return count;
7898            }
7899        }
7900
7901        if (*step < params->lbfgs.min_step) {
7902            return GGML_LINESEARCH_MINIMUM_STEP;
7903        }
7904        if (*step > params->lbfgs.max_step) {
7905            return GGML_LINESEARCH_MAXIMUM_STEP;
7906        }
7907        if (params->lbfgs.max_linesearch <= count) {
7908            return GGML_LINESEARCH_MAXIMUM_ITERATIONS;
7909        }
7910
7911        (*step) *= width;
7912    }
7913
7914    return GGML_LINESEARCH_FAIL;
7915}
7916
7917static enum ggml_opt_result ggml_opt_lbfgs(
7918        struct ggml_context * ctx,
7919        struct ggml_opt_params params,
7920        struct ggml_tensor * f,
7921        struct ggml_cgraph * gf,
7922        struct ggml_cgraph * gb) {
7923    if (params.lbfgs.linesearch == GGML_LINESEARCH_BACKTRACKING_WOLFE ||
7924        params.lbfgs.linesearch == GGML_LINESEARCH_BACKTRACKING_STRONG_WOLFE) {
7925        if (params.lbfgs.wolfe <= params.lbfgs.ftol || 1. <= params.lbfgs.wolfe) {
7926            return GGML_OPT_INVALID_WOLFE;
7927        }
7928    }
7929
7930    gf->n_threads = params.n_threads;
7931    gb->n_threads = params.n_threads;
7932
7933    const int m = params.lbfgs.m;
7934
7935    // these will store the parameters we want to optimize
7936    struct ggml_tensor * ps[GGML_MAX_PARAMS];
7937
7938    int np = 0;
7939    int nx = 0;
7940    for (int i = 0; i < gf->n_nodes; ++i) {
7941        if (gf->nodes[i]->is_param) {
7942            GGML_PRINT_DEBUG("found param %d: grad->op = %d\n", np, gf->nodes[i]->grad->op);
7943
7944            assert(np < GGML_MAX_PARAMS);
7945
7946            ps[np++] = gf->nodes[i];
7947            nx += ggml_nelements(gf->nodes[i]);
7948        }
7949    }
7950
7951    float * x  = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data; // current parameters
7952    float * xp = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data; // previous parameters
7953    float * g  = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data; // current gradient
7954    float * gp = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data; // previous gradient
7955    float * d  = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data; // search direction
7956
7957    float * pf = params.past > 0 ? ggml_new_tensor_1d(ctx, GGML_TYPE_F32, params.past)->data : NULL; // past function values
7958
7959    float fx    = 0.0f; // cost function value
7960    float xnorm = 0.0f; // ||x||
7961    float gnorm = 0.0f; // ||g||
7962    float step  = 0.0f;
7963
7964    // initialize x from the graph nodes
7965    ggml_opt_get_params(np, ps, x);
7966
7967    // the L-BFGS memory
7968    struct ggml_lbfgs_iteration_data * lm = alloca(sizeof(struct ggml_lbfgs_iteration_data)*m);
7969
7970    for (int i = 0; i < m; ++i) {
7971        lm[i].alpha = 0.0f;
7972        lm[i].ys    = 0.0f;
7973        lm[i].s     = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data;
7974        lm[i].y     = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nx)->data;
7975    }
7976
7977    // evaluate the function value and its gradient
7978    {
7979        ggml_opt_set_params(np, ps, x);
7980
7981        ggml_graph_reset  (gf);
7982        ggml_set_f32      (f->grad, 1.0f);
7983        ggml_graph_compute(ctx, gb);
7984
7985        ggml_opt_get_grad(np, ps, g);
7986
7987        fx = ggml_get_f32_1d(f, 0);
7988    }
7989
7990    if (pf) {
7991        pf[0] = fx;
7992    }
7993
7994    float fx_best = fx;
7995
7996    // search direction = -gradient
7997    ggml_vec_neg_f32(nx, d, g);
7998
7999    // ||x||, ||g||
8000    ggml_vec_norm_f32(nx, &xnorm, x);
8001    ggml_vec_norm_f32(nx, &gnorm, g);
8002
8003    if (xnorm < 1.0f) {
8004        xnorm = 1.0f;
8005    }
8006
8007    // already optimized
8008    if (gnorm/xnorm <= params.lbfgs.eps) {
8009        return GGML_OPT_OK;
8010    }
8011
8012    // initial step
8013    ggml_vec_norm_inv_f32(nx, &step, d);
8014
8015    int j                = 0;
8016    int k                = 1;
8017    int ls               = 0;
8018    int end              = 0;
8019    int bound            = 0;
8020    int n_no_improvement = 0;
8021
8022    float ys   = 0.0f;
8023    float yy   = 0.0f;
8024    float beta = 0.0f;
8025
8026    while (true) {
8027        // store the current position and gradient vectors
8028        ggml_vec_cpy_f32(nx, xp, x);
8029        ggml_vec_cpy_f32(nx, gp, g);
8030
8031        ls = linesearch_backtracking(ctx, &params, nx, x, &fx, g, d, &step, xp, f, gf, gb, np, ps);
8032
8033        if (ls < 0) {
8034            // linesearch failed - go back to the previous point and return
8035            ggml_vec_cpy_f32(nx, x, xp);
8036            ggml_vec_cpy_f32(nx, g, gp);
8037
8038            return ls;
8039        }
8040
8041        ggml_vec_norm_f32(nx, &xnorm, x);
8042        ggml_vec_norm_f32(nx, &gnorm, g);
8043
8044        GGML_PRINT_DEBUG("f = %10.6f\n", ggml_get_f32_1d(f, 0));
8045
8046        if (xnorm < 1.0) {
8047            xnorm = 1.0;
8048        }
8049        if (gnorm/xnorm <= params.lbfgs.eps) {
8050            // converged
8051            return GGML_OPT_OK;
8052        }
8053
8054        // delta-based convergence test
8055        if (pf != NULL) {
8056            // need at least params.past iterations to start checking for convergence
8057            if (params.past <= k) {
8058                const float rate = (pf[k%params.past] - fx)/fx;
8059
8060                if (fabs(rate) < params.delta) {
8061                    return GGML_OPT_OK;
8062                }
8063            }
8064
8065            pf[k%params.past] = fx;
8066        }
8067
8068        // check for improvement
8069        if (params.max_no_improvement > 0) {
8070            if (fx < fx_best) {
8071                fx_best = fx;
8072                n_no_improvement = 0;
8073            } else {
8074                n_no_improvement++;
8075
8076                if (n_no_improvement >= params.max_no_improvement) {
8077                    return GGML_OPT_OK;
8078                }
8079            }
8080        }
8081
8082        if (params.lbfgs.n_iter != 0 && params.lbfgs.n_iter < k + 1) {
8083            // reached the maximum number of iterations
8084            return GGML_OPT_DID_NOT_CONVERGE;
8085        }
8086
8087        // update vectors s and y:
8088        //   s_{k+1} = x_{k+1} - x_{k} = \step * d_{k}.
8089        //   y_{k+1} = g_{k+1} - g_{k}.
8090        //
8091        ggml_vec_sub_f32(nx, lm[end].s, x, xp);
8092        ggml_vec_sub_f32(nx, lm[end].y, g, gp);
8093
8094        // compute scalars ys and yy:
8095        //     ys = y^t \cdot s    -> 1 / \rho.
8096        //     yy = y^t \cdot y.
8097        //
8098        ggml_vec_dot_f32(nx, &ys, lm[end].y, lm[end].s);
8099        ggml_vec_dot_f32(nx, &yy, lm[end].y, lm[end].y);
8100
8101        lm[end].ys = ys;
8102
8103        // find new search direction
8104        //   ref: https://en.wikipedia.org/wiki/Limited-memory_BFGS
8105
8106        bound = (m <= k) ? m : k;
8107        k++;
8108        end = (end + 1)%m;
8109
8110        // initialize search direction with -g
8111        ggml_vec_neg_f32(nx, d, g);
8112
8113        j = end;
8114        for (int i = 0; i < bound; ++i) {
8115            j = (j + m - 1) % m;
8116            // \alpha_{j} = \rho_{j} s^{t}_{j} \cdot q_{k+1}
8117            ggml_vec_dot_f32(nx, &lm[j].alpha, lm[j].s, d);
8118            lm[j].alpha /= lm[j].ys;
8119            // q_{i} = q_{i+1} - \alpha_{i} y_{i}
8120            ggml_vec_mad_f32(nx, d, lm[j].y, -lm[j].alpha);
8121        }
8122
8123        ggml_vec_scale_f32(nx, d, ys/yy);
8124
8125        for (int i = 0; i < bound; ++i) {
8126            // \beta_{j} = \rho_{j} y^t_{j} \cdot \gamma_{i}
8127            ggml_vec_dot_f32(nx, &beta, lm[j].y, d);
8128            beta /= lm[j].ys;
8129            // \gamma_{i+1} = \gamma_{i} + (\alpha_{j} - \beta_{j}) s_{j}
8130            ggml_vec_mad_f32(nx, d, lm[j].s, lm[j].alpha - beta);
8131            j = (j + 1)%m;
8132        }
8133
8134        step = 1.0;
8135    }
8136
8137    return GGML_OPT_DID_NOT_CONVERGE;
8138}
8139
8140struct ggml_opt_params ggml_opt_default_params(enum ggml_opt_type type) {
8141    struct ggml_opt_params result;
8142
8143    switch (type) {
8144        case GGML_OPT_ADAM:
8145            {
8146                result = (struct ggml_opt_params) {
8147                    .type      = GGML_OPT_ADAM,
8148                    .n_threads = 1,
8149                    .past      = 0,
8150                    .delta     = 1e-5f,
8151
8152                    .max_no_improvement = 100,
8153
8154                    .print_forward_graph  = true,
8155                    .print_backward_graph = true,
8156
8157                    .adam = {
8158                        .n_iter = 10000,
8159                        .alpha  = 0.001f,
8160                        .beta1  = 0.9f,
8161                        .beta2  = 0.999f,
8162                        .eps    = 1e-8f,
8163                        .eps_f  = 1e-5f,
8164                        .eps_g  = 1e-3f,
8165                    },
8166                };
8167            } break;
8168        case GGML_OPT_LBFGS:
8169            {
8170                result = (struct ggml_opt_params) {
8171                    .type      = GGML_OPT_LBFGS,
8172                    .n_threads = 1,
8173                    .past      = 0,
8174                    .delta     = 1e-5f,
8175
8176                    .max_no_improvement = 0,
8177
8178                    .print_forward_graph  = true,
8179                    .print_backward_graph = true,
8180
8181                    .lbfgs = {
8182                        .m              = 6,
8183                        .n_iter         = 100,
8184                        .max_linesearch = 20,
8185
8186                        .eps      = 1e-5f,
8187                        .ftol     = 1e-4f,
8188                        .wolfe    = 0.9f,
8189                        .min_step = 1e-20f,
8190                        .max_step = 1e+20f,
8191
8192                        .linesearch = GGML_LINESEARCH_DEFAULT,
8193                    },
8194                };
8195            } break;
8196    }
8197
8198    return result;
8199}
8200
8201enum ggml_opt_result ggml_opt(
8202        struct ggml_context * ctx,
8203        struct ggml_opt_params params,
8204        struct ggml_tensor * f) {
8205    bool free_ctx = false;
8206    if (ctx == NULL) {
8207        struct ggml_init_params params_ctx = {
8208            .mem_size   = 16*1024*1024,
8209            .mem_buffer = NULL,
8210        };
8211
8212        ctx = ggml_init(params_ctx);
8213        if (ctx == NULL) {
8214            return GGML_OPT_NO_CONTEXT;
8215        }
8216
8217        free_ctx = true;
8218    }
8219
8220    enum ggml_opt_result result = GGML_OPT_OK;
8221
8222    // build forward + backward compute graphs
8223    struct ggml_cgraph gf = ggml_build_forward (f);
8224    struct ggml_cgraph gb = ggml_build_backward(ctx, &gf, false);
8225
8226    switch (params.type) {
8227        case GGML_OPT_ADAM:
8228            {
8229                result = ggml_opt_adam(ctx, params, f, &gf, &gb);
8230            } break;
8231        case GGML_OPT_LBFGS:
8232            {
8233                result = ggml_opt_lbfgs(ctx, params, f, &gf, &gb);
8234            } break;
8235    }
8236
8237    if (params.print_forward_graph) {
8238        ggml_graph_print   (&gf);
8239        ggml_graph_dump_dot(&gf, NULL, "opt-forward.dot");
8240    }
8241
8242    if (params.print_backward_graph) {
8243        ggml_graph_print   (&gb);
8244        ggml_graph_dump_dot(&gb, &gf, "opt-backward.dot");
8245    }
8246
8247    if (free_ctx) {
8248        ggml_free(ctx);
8249    }
8250
8251    return result;
8252}
8253
8254////////////////////////////////////////////////////////////////////////////////
8255
8256int ggml_cpu_has_avx(void) {
8257#if defined(__AVX__)
8258    return 1;
8259#else
8260    return 0;
8261#endif
8262}
8263
8264int ggml_cpu_has_avx2(void) {
8265#if defined(__AVX2__)
8266    return 1;
8267#else
8268    return 0;
8269#endif
8270}
8271
8272int ggml_cpu_has_avx512(void) {
8273#if defined(__AVX512F__)
8274    return 1;
8275#else
8276    return 0;
8277#endif
8278}
8279
8280int ggml_cpu_has_fma(void) {
8281#if defined(__FMA__)
8282    return 1;
8283#else
8284    return 0;
8285#endif
8286}
8287
8288int ggml_cpu_has_neon(void) {
8289#if defined(__ARM_NEON)
8290    return 1;
8291#else
8292    return 0;
8293#endif
8294}
8295
8296int ggml_cpu_has_arm_fma(void) {
8297#if defined(__ARM_FEATURE_FMA)
8298    return 1;
8299#else
8300    return 0;
8301#endif
8302}
8303
8304int ggml_cpu_has_f16c(void) {
8305#if defined(__F16C__)
8306    return 1;
8307#else
8308    return 0;
8309#endif
8310}
8311
8312int ggml_cpu_has_fp16_va(void) {
8313#if defined(__ARM_FEATURE_FP16_VECTOR_ARITHMETIC)
8314    return 1;
8315#else
8316    return 0;
8317#endif
8318}
8319
8320int ggml_cpu_has_wasm_simd(void) {
8321#if defined(__wasm_simd128__)
8322    return 1;
8323#else
8324    return 0;
8325#endif
8326}
8327
8328int ggml_cpu_has_blas(void) {
8329#if defined(GGML_USE_ACCELERATE) || defined(GGML_USE_OPENBLAS)
8330    return 1;
8331#else
8332    return 0;
8333#endif
8334}
8335
8336////////////////////////////////////////////////////////////////////////////////