Skip to content

File syn_matrix.c

File List > src > syntropic > util > syn_matrix.c

Go to the documentation of this file

#if __has_include("syn_config.h")
#include "syn_config.h"
#endif

#if !defined(SYN_USE_MATRIX) || SYN_USE_MATRIX

#include "../util/syn_assert.h"
#include "syn_matrix.h"

#include <string.h>

/* ════════════════════════════════════════════════════════════════════════ */
/*  Core operations                                                        */
/* ════════════════════════════════════════════════════════════════════════ */

void syn_matrix_identity(SYN_Matrix *m)
{
    SYN_ASSERT(m != NULL);
    SYN_ASSERT(m->rows == m->cols); /* Must be square */

    memset(m->data, 0, (size_t)m->rows * m->cols * sizeof(q16_t));
    uint8_t i;
    for (i = 0; i < m->rows; i++) {
        SYN_MAT_AT(m, i, i) = Q16_ONE;
    }
}

void syn_matrix_zero(SYN_Matrix *m)
{
    SYN_ASSERT(m != NULL);
    memset(m->data, 0, (size_t)m->rows * m->cols * sizeof(q16_t));
}

void syn_matrix_copy(SYN_Matrix *dst, const SYN_Matrix *src)
{
    SYN_ASSERT(dst != NULL && src != NULL);
    SYN_ASSERT(dst->rows == src->rows && dst->cols == src->cols);
    memcpy(dst->data, src->data, (size_t)src->rows * src->cols * sizeof(q16_t));
}

void syn_matrix_add(const SYN_Matrix *a, const SYN_Matrix *b, SYN_Matrix *out)
{
    SYN_ASSERT(a != NULL && b != NULL && out != NULL);
    SYN_ASSERT(a->rows == b->rows && a->cols == b->cols);
    SYN_ASSERT(out->rows == a->rows && out->cols == a->cols);

    uint16_t n = (uint16_t)a->rows * a->cols;
    uint16_t i;
    for (i = 0; i < n; i++) {
        out->data[i] = a->data[i] + b->data[i];
    }
}

void syn_matrix_sub(const SYN_Matrix *a, const SYN_Matrix *b, SYN_Matrix *out)
{
    SYN_ASSERT(a != NULL && b != NULL && out != NULL);
    if (a == NULL || b == NULL || out == NULL)
        return;
    SYN_ASSERT(a->rows == b->rows && a->cols == b->cols);
    SYN_ASSERT(out->rows == a->rows && out->cols == a->cols);

    uint16_t n = (uint16_t)a->rows * a->cols;
    uint16_t i;
    for (i = 0; i < n; i++) {
        out->data[i] = a->data[i] - b->data[i];
    }
}

void syn_matrix_scale(const SYN_Matrix *a, q16_t scalar, SYN_Matrix *out)
{
    SYN_ASSERT(a != NULL && out != NULL);
    SYN_ASSERT(out->rows == a->rows && out->cols == a->cols);

    uint16_t n = (uint16_t)a->rows * a->cols;
    uint16_t i;
    for (i = 0; i < n; i++) {
        out->data[i] = q16_mul(a->data[i], scalar);
    }
}

/* ════════════════════════════════════════════════════════════════════════ */
/*  Multiply                                                               */
/* ════════════════════════════════════════════════════════════════════════ */

void syn_matrix_mul(const SYN_Matrix *a, const SYN_Matrix *b, SYN_Matrix *out)
{
    SYN_ASSERT(a != NULL && b != NULL && out != NULL);
    SYN_ASSERT(a->cols == b->rows);
    SYN_ASSERT(out->rows == a->rows && out->cols == b->cols);

    uint8_t i, j, k;

    /* Unrolled fast-paths for common control/DSP matrix sizes */
    if (a->rows == 2 && a->cols == 2 && b->cols == 2) {
        q16_t a00 = a->data[0], a01 = a->data[1], a10 = a->data[2], a11 = a->data[3];
        q16_t b00 = b->data[0], b01 = b->data[1], b10 = b->data[2], b11 = b->data[3];
        out->data[0] = (q16_t)(((int64_t)a00 * b00 + (int64_t)a01 * b10) >> Q16_SHIFT);
        out->data[1] = (q16_t)(((int64_t)a00 * b01 + (int64_t)a01 * b11) >> Q16_SHIFT);
        out->data[2] = (q16_t)(((int64_t)a10 * b00 + (int64_t)a11 * b10) >> Q16_SHIFT);
        out->data[3] = (q16_t)(((int64_t)a10 * b01 + (int64_t)a11 * b11) >> Q16_SHIFT);
        return;
    }

    if (a->rows == 3 && a->cols == 3 && b->cols == 3) {
        const q16_t *ad = a->data;
        const q16_t *bd = b->data;
        q16_t *od = out->data;
        for (i = 0; i < 3; i++) {
            q16_t a0 = ad[i * 3], a1 = ad[i * 3 + 1], a2 = ad[i * 3 + 2];
            od[i * 3] = (q16_t)(((int64_t)a0 * bd[0] + (int64_t)a1 * bd[3] + (int64_t)a2 * bd[6]) >>
                                Q16_SHIFT);
            od[i * 3 + 1] =
                (q16_t)(((int64_t)a0 * bd[1] + (int64_t)a1 * bd[4] + (int64_t)a2 * bd[7]) >>
                        Q16_SHIFT);
            od[i * 3 + 2] =
                (q16_t)(((int64_t)a0 * bd[2] + (int64_t)a1 * bd[5] + (int64_t)a2 * bd[8]) >>
                        Q16_SHIFT);
        }
        return;
    }

    if (a->rows == 4 && a->cols == 4 && b->cols == 4) {
        const q16_t *ad = a->data;
        const q16_t *bd = b->data;
        q16_t *od = out->data;
        for (i = 0; i < 4; i++) {
            q16_t a0 = ad[i * 4], a1 = ad[i * 4 + 1], a2 = ad[i * 4 + 2], a3 = ad[i * 4 + 3];
            od[i * 4] = (q16_t)(((int64_t)a0 * bd[0] + (int64_t)a1 * bd[4] + (int64_t)a2 * bd[8] +
                                 (int64_t)a3 * bd[12]) >>
                                Q16_SHIFT);
            od[i * 4 + 1] = (q16_t)(((int64_t)a0 * bd[1] + (int64_t)a1 * bd[5] +
                                     (int64_t)a2 * bd[9] + (int64_t)a3 * bd[13]) >>
                                    Q16_SHIFT);
            od[i * 4 + 2] = (q16_t)(((int64_t)a0 * bd[2] + (int64_t)a1 * bd[6] +
                                     (int64_t)a2 * bd[10] + (int64_t)a3 * bd[14]) >>
                                    Q16_SHIFT);
            od[i * 4 + 3] = (q16_t)(((int64_t)a0 * bd[3] + (int64_t)a1 * bd[7] +
                                     (int64_t)a2 * bd[11] + (int64_t)a3 * bd[15]) >>
                                    Q16_SHIFT);
        }
        return;
    }

    /* General M×N×K matrix multiplication */
    for (i = 0; i < a->rows; i++) {
        for (j = 0; j < b->cols; j++) {
            int64_t acc = 0;
            for (k = 0; k < a->cols; k++) {
                acc += (int64_t)SYN_MAT_AT(a, i, k) * SYN_MAT_AT(b, k, j);
            }
            SYN_MAT_AT(out, i, j) = (q16_t)(acc >> Q16_SHIFT);
        }
    }
}

void syn_matrix_mul_vec(const SYN_Matrix *m, const q16_t *v_in, q16_t *v_out, uint8_t n_in)
{
    SYN_ASSERT(m != NULL && v_in != NULL && v_out != NULL);
    SYN_ASSERT(m->cols == n_in);

    uint8_t i, k;
    for (i = 0; i < m->rows; i++) {
        int64_t acc = 0;
        for (k = 0; k < n_in; k++) {
            acc += (int64_t)SYN_MAT_AT(m, i, k) * v_in[k];
        }
        v_out[i] = (q16_t)(acc >> Q16_SHIFT);
    }
}

/* ════════════════════════════════════════════════════════════════════════ */
/*  Transpose, trace                                                       */
/* ════════════════════════════════════════════════════════════════════════ */

void syn_matrix_transpose(const SYN_Matrix *a, SYN_Matrix *out)
{
    SYN_ASSERT(a != NULL && out != NULL);
    SYN_ASSERT(out->rows == a->cols && out->cols == a->rows);

    uint8_t i, j;
    for (i = 0; i < a->rows; i++) {
        for (j = 0; j < a->cols; j++) {
            SYN_MAT_AT(out, j, i) = SYN_MAT_AT(a, i, j);
        }
    }
}

q16_t syn_matrix_trace(const SYN_Matrix *m)
{
    SYN_ASSERT(m != NULL);
    SYN_ASSERT(m->rows == m->cols);

    q16_t sum = 0;
    uint8_t i;
    for (i = 0; i < m->rows; i++) {
        sum += SYN_MAT_AT(m, i, i);
    }
    return sum;
}

/* ════════════════════════════════════════════════════════════════════════ */
/*  Determinant                                                            */
/* ════════════════════════════════════════════════════════════════════════ */

static q16_t det_2x2(const q16_t *d)
{
    return (q16_t)(((int64_t)d[0] * d[3] - (int64_t)d[1] * d[2]) >> Q16_SHIFT);
}

static q16_t det_3x3(const q16_t *d)
{
    /* a(ei - fh) - b(di - fg) + c(dh - eg) */
    int64_t a = d[0], b = d[1], c = d[2];
    int64_t det_a = ((int64_t)d[4] * d[8] - (int64_t)d[5] * d[7]) >> Q16_SHIFT;
    int64_t det_b = ((int64_t)d[3] * d[8] - (int64_t)d[5] * d[6]) >> Q16_SHIFT;
    int64_t det_c = ((int64_t)d[3] * d[7] - (int64_t)d[4] * d[6]) >> Q16_SHIFT;

    int64_t result = (a * det_a - b * det_b + c * det_c) >> Q16_SHIFT;
    return (q16_t)result;
}

static q16_t det_4x4(const q16_t *d)
{
    /* Cofactor expansion along first row */
    q16_t minor0[9], minor1[9], minor2[9], minor3[9];

    /* Minor of d[0]: rows 1-3, cols 1-3 */
    minor0[0] = d[5];
    minor0[1] = d[6];
    minor0[2] = d[7];
    minor0[3] = d[9];
    minor0[4] = d[10];
    minor0[5] = d[11];
    minor0[6] = d[13];
    minor0[7] = d[14];
    minor0[8] = d[15];

    /* Minor of d[1]: rows 1-3, cols 0,2,3 */
    minor1[0] = d[4];
    minor1[1] = d[6];
    minor1[2] = d[7];
    minor1[3] = d[8];
    minor1[4] = d[10];
    minor1[5] = d[11];
    minor1[6] = d[12];
    minor1[7] = d[14];
    minor1[8] = d[15];

    /* Minor of d[2]: rows 1-3, cols 0,1,3 */
    minor2[0] = d[4];
    minor2[1] = d[5];
    minor2[2] = d[7];
    minor2[3] = d[8];
    minor2[4] = d[9];
    minor2[5] = d[11];
    minor2[6] = d[12];
    minor2[7] = d[13];
    minor2[8] = d[15];

    /* Minor of d[3]: rows 1-3, cols 0,1,2 */
    minor3[0] = d[4];
    minor3[1] = d[5];
    minor3[2] = d[6];
    minor3[3] = d[8];
    minor3[4] = d[9];
    minor3[5] = d[10];
    minor3[6] = d[12];
    minor3[7] = d[13];
    minor3[8] = d[14];

    int64_t result = ((int64_t)d[0] * det_3x3(minor0)) >> Q16_SHIFT;
    result -= ((int64_t)d[1] * det_3x3(minor1)) >> Q16_SHIFT;
    result += ((int64_t)d[2] * det_3x3(minor2)) >> Q16_SHIFT;
    result -= ((int64_t)d[3] * det_3x3(minor3)) >> Q16_SHIFT;

    return (q16_t)result;
}

q16_t syn_matrix_det(const SYN_Matrix *m)
{
    SYN_ASSERT(m != NULL);
    SYN_ASSERT(m->rows == m->cols);

    switch (m->rows) {
    case 1:
        return m->data[0];
    case 2:
        return det_2x2(m->data);
    case 3:
        return det_3x3(m->data);
    case 4:
        return det_4x4(m->data);
    default:
        return 0; /* Unsupported */
    }
}

/* ════════════════════════════════════════════════════════════════════════ */
/*  Inverse                                                                */
/* ════════════════════════════════════════════════════════════════════════ */

static SYN_Status inv_2x2(const SYN_Matrix *m, SYN_Matrix *out)
{
    q16_t det = det_2x2(m->data);
    if (det == 0)
        return SYN_ERROR;

    /* [a b]^-1 = (1/det) * [ d -b]
     * [c d]                [-c  a] */
    SYN_MAT_AT(out, 0, 0) = q16_div(m->data[3], det);
    SYN_MAT_AT(out, 0, 1) = -q16_div(m->data[1], det);
    SYN_MAT_AT(out, 1, 0) = -q16_div(m->data[2], det);
    SYN_MAT_AT(out, 1, 1) = q16_div(m->data[0], det);

    return SYN_OK;
}

static SYN_Status inv_3x3(const SYN_Matrix *m, SYN_Matrix *out)
{
    q16_t det = det_3x3(m->data);
    if (det == 0)
        return SYN_ERROR;

    const q16_t *d = m->data;

    /* Cofactor matrix (transposed = adjugate) divided by det */
    /* Row 0 of adjugate (cofactors of column 0) */
    q16_t c00 = (q16_t)(((int64_t)d[4] * d[8] - (int64_t)d[5] * d[7]) >> Q16_SHIFT);
    q16_t c01 = (q16_t)(((int64_t)d[2] * d[7] - (int64_t)d[1] * d[8]) >> Q16_SHIFT);
    q16_t c02 = (q16_t)(((int64_t)d[1] * d[5] - (int64_t)d[2] * d[4]) >> Q16_SHIFT);

    q16_t c10 = (q16_t)(((int64_t)d[5] * d[6] - (int64_t)d[3] * d[8]) >> Q16_SHIFT);
    q16_t c11 = (q16_t)(((int64_t)d[0] * d[8] - (int64_t)d[2] * d[6]) >> Q16_SHIFT);
    q16_t c12 = (q16_t)(((int64_t)d[2] * d[3] - (int64_t)d[0] * d[5]) >> Q16_SHIFT);

    q16_t c20 = (q16_t)(((int64_t)d[3] * d[7] - (int64_t)d[4] * d[6]) >> Q16_SHIFT);
    q16_t c21 = (q16_t)(((int64_t)d[1] * d[6] - (int64_t)d[0] * d[7]) >> Q16_SHIFT);
    q16_t c22 = (q16_t)(((int64_t)d[0] * d[4] - (int64_t)d[1] * d[3]) >> Q16_SHIFT);

    SYN_MAT_AT(out, 0, 0) = q16_div(c00, det);
    SYN_MAT_AT(out, 0, 1) = q16_div(c01, det);
    SYN_MAT_AT(out, 0, 2) = q16_div(c02, det);
    SYN_MAT_AT(out, 1, 0) = q16_div(c10, det);
    SYN_MAT_AT(out, 1, 1) = q16_div(c11, det);
    SYN_MAT_AT(out, 1, 2) = q16_div(c12, det);
    SYN_MAT_AT(out, 2, 0) = q16_div(c20, det);
    SYN_MAT_AT(out, 2, 1) = q16_div(c21, det);
    SYN_MAT_AT(out, 2, 2) = q16_div(c22, det);

    return SYN_OK;
}

static SYN_Status inv_4x4(const SYN_Matrix *m, SYN_Matrix *out)
{
    /* Work on a copy augmented with identity: [M | I] */
    q16_t aug[4][8];
    uint8_t i, j, k;

    /* Initialize augmented matrix */
    for (i = 0; i < 4; i++) {
        for (j = 0; j < 4; j++) {
            aug[i][j] = SYN_MAT_AT(m, i, j);
            aug[i][j + 4] = (i == j) ? Q16_ONE : 0;
        }
    }

    /* Gauss-Jordan elimination with partial pivoting */
    for (k = 0; k < 4; k++) {
        /* Find pivot */
        uint8_t max_row = k;
        q16_t max_val = q16_abs(aug[k][k]);
        for (i = k + 1; i < 4; i++) {
            q16_t val = q16_abs(aug[i][k]);
            if (val > max_val) {
                max_val = val;
                max_row = i;
            }
        }

        if (max_val == 0)
            return SYN_ERROR; /* Singular */

        /* Swap rows */
        if (max_row != k) {
            for (j = 0; j < 8; j++) {
                q16_t tmp = aug[k][j];
                aug[k][j] = aug[max_row][j];
                aug[max_row][j] = tmp;
            }
        }

        /* Scale pivot row so aug[k][k] = 1.0 */
        q16_t pivot = aug[k][k];
        for (j = 0; j < 8; j++) {
            aug[k][j] = q16_div(aug[k][j], pivot);
        }

        /* Eliminate column k from all other rows */
        for (i = 0; i < 4; i++) {
            if (i == k)
                continue;
            q16_t factor = aug[i][k];
            for (j = 0; j < 8; j++) {
                aug[i][j] -= q16_mul(factor, aug[k][j]);
            }
        }
    }

    /* Extract inverse from right half of augmented matrix */
    for (i = 0; i < 4; i++) {
        for (j = 0; j < 4; j++) {
            SYN_MAT_AT(out, i, j) = aug[i][j + 4];
        }
    }

    return SYN_OK;
}

static SYN_Status inv_1x1(const SYN_Matrix *m, SYN_Matrix *out)
{
    if (m->data[0] == 0)
        return SYN_ERROR;
    out->data[0] = q16_div(Q16_ONE, m->data[0]);
    return SYN_OK;
}

SYN_Status syn_matrix_inv(const SYN_Matrix *m, SYN_Matrix *out)
{
    SYN_ASSERT(m != NULL && out != NULL);
    SYN_ASSERT(m->rows == m->cols);
    SYN_ASSERT(out->rows == m->rows && out->cols == m->cols);

    switch (m->rows) {
    case 1:
        return inv_1x1(m, out);
    case 2:
        return inv_2x2(m, out);
    case 3:
        return inv_3x3(m, out);
    case 4:
        return inv_4x4(m, out);
    default:
        return syn_matrix_inv_lu(m, out);
    }
}

SYN_Status syn_matrix_inv_lu_work(const SYN_Matrix *src, SYN_Matrix *dst, q16_t *lu_work,
                                  uint8_t *p_work, q16_t *y_work)
{
    SYN_ASSERT(src != NULL && dst != NULL);
    SYN_ASSERT(lu_work != NULL && p_work != NULL && y_work != NULL);

    uint8_t n = src->rows;
    if (src->cols != n || dst->rows != n || dst->cols != n || n > SYN_SOLVER_MAX_N) {
        return SYN_INVALID_PARAM;
    }

    for (uint8_t j = 0; j < n; j++) {
        q16_t ej_data[SYN_SOLVER_MAX_N];
        for (uint8_t i = 0; i < n; i++) {
            ej_data[i] = (i == j) ? Q16_ONE : 0;
        }
        SYN_Matrix ej = {ej_data, n, 1};

        q16_t x_data[SYN_SOLVER_MAX_N];
        SYN_Matrix xj = {x_data, n, 1};

        SYN_Status status = syn_matrix_solve_lu_work(src, &ej, &xj, lu_work, p_work, y_work);
        if (status != SYN_OK) {
            return status;
        }

        for (uint8_t i = 0; i < n; i++) {
            SYN_MAT_AT(dst, i, j) = xj.data[i];
        }
    }

    return SYN_OK;
}

SYN_Status syn_matrix_inv_lu(const SYN_Matrix *src, SYN_Matrix *dst)
{
    SYN_ASSERT(src != NULL && dst != NULL);
    uint8_t n = src->rows;
    if (src->cols != n || dst->rows != n || dst->cols != n || n > SYN_SOLVER_MAX_N) {
        return SYN_INVALID_PARAM;
    }
    q16_t lu_work[n * n];
    uint8_t p_work[n];
    q16_t col_work[n];
    return syn_matrix_inv_lu_work(src, dst, lu_work, p_work, col_work);
}

/* ════════════════════════════════════════════════════════════════════════ */
/*  2D transforms (3×3 homogeneous)                                        */
/* ════════════════════════════════════════════════════════════════════════ */

void syn_matrix_rotate_2d(SYN_Matrix *out, q16_t angle)
{
    SYN_ASSERT(out != NULL);
    SYN_ASSERT(out->rows == 3 && out->cols == 3);

    q16_t c = q16_cos(angle);
    q16_t s = q16_sin(angle);

    syn_matrix_identity(out);
    SYN_MAT_AT(out, 0, 0) = c;
    SYN_MAT_AT(out, 0, 1) = -s;
    SYN_MAT_AT(out, 1, 0) = s;
    SYN_MAT_AT(out, 1, 1) = c;
}

void syn_matrix_translate_2d(SYN_Matrix *out, q16_t tx, q16_t ty)
{
    SYN_ASSERT(out != NULL);
    SYN_ASSERT(out->rows == 3 && out->cols == 3);

    syn_matrix_identity(out);
    SYN_MAT_AT(out, 0, 2) = tx;
    SYN_MAT_AT(out, 1, 2) = ty;
}

void syn_matrix_scale_2d(SYN_Matrix *out, q16_t sx, q16_t sy)
{
    SYN_ASSERT(out != NULL);
    SYN_ASSERT(out->rows == 3 && out->cols == 3);

    syn_matrix_zero(out);
    SYN_MAT_AT(out, 0, 0) = sx;
    SYN_MAT_AT(out, 1, 1) = sy;
    SYN_MAT_AT(out, 2, 2) = Q16_ONE;
}

/* ════════════════════════════════════════════════════════════════════════ */
/*  3D transforms (4×4 homogeneous)                                        */
/* ════════════════════════════════════════════════════════════════════════ */

void syn_matrix_rotate_x(SYN_Matrix *out, q16_t angle)
{
    SYN_ASSERT(out != NULL);
    SYN_ASSERT(out->rows == 4 && out->cols == 4);

    q16_t c = q16_cos(angle);
    q16_t s = q16_sin(angle);

    syn_matrix_identity(out);
    SYN_MAT_AT(out, 1, 1) = c;
    SYN_MAT_AT(out, 1, 2) = -s;
    SYN_MAT_AT(out, 2, 1) = s;
    SYN_MAT_AT(out, 2, 2) = c;
}

void syn_matrix_rotate_y(SYN_Matrix *out, q16_t angle)
{
    SYN_ASSERT(out != NULL);
    SYN_ASSERT(out->rows == 4 && out->cols == 4);

    q16_t c = q16_cos(angle);
    q16_t s = q16_sin(angle);

    syn_matrix_identity(out);
    SYN_MAT_AT(out, 0, 0) = c;
    SYN_MAT_AT(out, 0, 2) = s;
    SYN_MAT_AT(out, 2, 0) = -s;
    SYN_MAT_AT(out, 2, 2) = c;
}

void syn_matrix_rotate_z(SYN_Matrix *out, q16_t angle)
{
    SYN_ASSERT(out != NULL);
    SYN_ASSERT(out->rows == 4 && out->cols == 4);

    q16_t c = q16_cos(angle);
    q16_t s = q16_sin(angle);

    syn_matrix_identity(out);
    SYN_MAT_AT(out, 0, 0) = c;
    SYN_MAT_AT(out, 0, 1) = -s;
    SYN_MAT_AT(out, 1, 0) = s;
    SYN_MAT_AT(out, 1, 1) = c;
}

void syn_matrix_translate_3d(SYN_Matrix *out, q16_t tx, q16_t ty, q16_t tz)
{
    SYN_ASSERT(out != NULL);
    SYN_ASSERT(out->rows == 4 && out->cols == 4);

    syn_matrix_identity(out);
    SYN_MAT_AT(out, 0, 3) = tx;
    SYN_MAT_AT(out, 1, 3) = ty;
    SYN_MAT_AT(out, 2, 3) = tz;
}

/* ════════════════════════════════════════════════════════════════════════ */
/*  Vector helpers                                                         */
/* ════════════════════════════════════════════════════════════════════════ */

q16_t syn_vec_dot(const q16_t *a, const q16_t *b, uint8_t n)
{
    SYN_ASSERT(a != NULL && b != NULL);

    int64_t acc = 0;
    uint8_t i;
    for (i = 0; i < n; i++) {
        acc += (int64_t)a[i] * b[i];
    }
    return (q16_t)(acc >> Q16_SHIFT);
}

void syn_vec3_cross(const q16_t *a, const q16_t *b, q16_t *out)
{
    SYN_ASSERT(a != NULL && b != NULL && out != NULL);

    out[0] = (q16_t)(((int64_t)a[1] * b[2] - (int64_t)a[2] * b[1]) >> Q16_SHIFT);
    out[1] = (q16_t)(((int64_t)a[2] * b[0] - (int64_t)a[0] * b[2]) >> Q16_SHIFT);
    out[2] = (q16_t)(((int64_t)a[0] * b[1] - (int64_t)a[1] * b[0]) >> Q16_SHIFT);
}

q16_t syn_vec_norm(const q16_t *v, uint8_t n)
{
    SYN_ASSERT(v != NULL);

    int64_t sum = 0;
    uint8_t i;
    for (i = 0; i < n; i++) {
        sum += (int64_t)v[i] * v[i];
    }

    /* sum is in Q32.32. Convert to Q16.16 for q16_sqrt. */
    q16_t sum_q16 = (q16_t)(sum >> Q16_SHIFT);
    return q16_sqrt(sum_q16);
}

SYN_Status syn_vec_normalize(const q16_t *v, q16_t *out, uint8_t n)
{
    SYN_ASSERT(v != NULL && out != NULL);

    q16_t mag = syn_vec_norm(v, n);
    if (mag == 0)
        return SYN_ERROR;

    uint8_t i;
    for (i = 0; i < n; i++) {
        out[i] = q16_div(v[i], mag);
    }
    return SYN_OK;
}

/* ── Linear Solvers ─────────────────────────────────────────────────────── */

SYN_Status syn_matrix_solve_lu_work(const SYN_Matrix *A, const SYN_Matrix *b, SYN_Matrix *x,
                                    q16_t *lu, uint8_t *P, q16_t *y)
{
    SYN_ASSERT(A != NULL && b != NULL && x != NULL);
    SYN_ASSERT(lu != NULL && P != NULL && y != NULL);

    uint8_t n = A->rows;
    if (A->cols != n || b->rows != n || b->cols != 1 || x->rows != n || x->cols != 1) {
        return SYN_INVALID_PARAM;
    }
    if (n > SYN_SOLVER_MAX_N)
        return SYN_INVALID_PARAM;

    uint8_t i, j, k;

    for (i = 0; i < n; i++) {
        P[i] = i;
        for (j = 0; j < n; j++) {
            lu[i * n + j] = SYN_MAT_AT(A, i, j);
        }
    }

    /* Doolittle LU factorization with partial pivoting */
    for (i = 0; i < n; i++) {
        /* Pivot selection */
        q16_t max_val = 0;
        uint8_t pivot_idx = i;
        for (j = i; j < n; j++) {
            q16_t val = q16_abs(lu[j * n + i]);
            if (val > max_val) {
                max_val = val;
                pivot_idx = j;
            }
        }

        if (max_val == 0)
            return SYN_ERROR; /* Singular matrix */

        /* Swap rows if needed */
        if (pivot_idx != i) {
            uint8_t tmp_p = P[i];
            P[i] = P[pivot_idx];
            P[pivot_idx] = tmp_p;
            for (j = 0; j < n; j++) {
                q16_t tmp_v = lu[i * n + j];
                lu[i * n + j] = lu[pivot_idx * n + j];
                lu[pivot_idx * n + j] = tmp_v;
            }
        }

        /* Elimination */
        for (j = i + 1; j < n; j++) {
            lu[j * n + i] = q16_div(lu[j * n + i], lu[i * n + i]);
            for (k = i + 1; k < n; k++) {
                int64_t mult = (int64_t)lu[j * n + i] * lu[i * n + k];
                lu[j * n + k] -= (q16_t)(mult >> Q16_SHIFT);
            }
        }
    }

    /* Forward substitution L · y = P · b */
    for (i = 0; i < n; i++) {
        int64_t sum = (int64_t)b->data[P[i]];
        for (j = 0; j < i; j++) {
            sum -= ((int64_t)lu[i * n + j] * y[j]) >> Q16_SHIFT;
        }
        y[i] = (q16_t)sum;
    }

    /* Back substitution U · x = y */
    for (i = n; i > 0; i--) {
        uint8_t idx = i - 1;
        int64_t sum = (int64_t)y[idx];
        for (j = idx + 1; j < n; j++) {
            sum -= ((int64_t)lu[idx * n + j] * x->data[j]) >> Q16_SHIFT;
        }
        x->data[idx] = q16_div((q16_t)sum, lu[idx * n + idx]);
    }

    return SYN_OK;
}

SYN_Status syn_matrix_solve_lu(const SYN_Matrix *A, const SYN_Matrix *b, SYN_Matrix *x)
{
    SYN_ASSERT(A != NULL && b != NULL && x != NULL);
    uint8_t n = A->rows;
    if (n > SYN_SOLVER_MAX_N)
        return SYN_INVALID_PARAM;
    q16_t lu[n * n];
    uint8_t P[n];
    q16_t y[n];
    return syn_matrix_solve_lu_work(A, b, x, lu, P, y);
}

SYN_Status syn_matrix_solve_cholesky_work(const SYN_Matrix *A, const SYN_Matrix *b, SYN_Matrix *x,
                                          q16_t *L, q16_t *y)
{
    SYN_ASSERT(A != NULL && b != NULL && x != NULL);
    SYN_ASSERT(L != NULL && y != NULL);

    uint8_t n = A->rows;
    if (A->cols != n || b->rows != n || b->cols != 1 || x->rows != n || x->cols != 1) {
        return SYN_INVALID_PARAM;
    }
    if (n == 0 || n > SYN_SOLVER_MAX_N)
        return SYN_INVALID_PARAM;

    memset(L, 0, (size_t)n * n * sizeof(q16_t));
    memset(y, 0, (size_t)n * sizeof(q16_t));

    uint8_t i, j, k;
    for (i = 0; i < n; i++) {
        for (j = 0; j <= i; j++) {
            int64_t sum = (int64_t)SYN_MAT_AT(A, i, j);
            for (k = 0; k < j; k++) {
                sum -= ((int64_t)L[i * n + k] * L[j * n + k]) >> Q16_SHIFT;
            }

            if (i == j) {
                if (sum <= 0)
                    return SYN_ERROR; /* Not positive-definite */
                L[i * n + j] = q16_sqrt((q16_t)sum);
            } else {
                L[i * n + j] = q16_div((q16_t)sum, L[j * n + j]);
            }
        }
    }

    /* Forward substitution L · y = b */
    for (i = 0; i < n; i++) {
        int64_t sum = (int64_t)b->data[i];
        for (j = 0; j < i; j++) {
            sum -= ((int64_t)L[i * n + j] * y[j]) >> Q16_SHIFT;
        }
        y[i] = q16_div((q16_t)sum, L[i * n + i]);
    }

    /* Back substitution Lᵀ · x = y */
    for (int idx = (int)n - 1; idx >= 0; idx--) {
        int64_t sum = (int64_t)y[idx];
        for (j = (uint8_t)(idx + 1); j < n; j++) {
            sum -= ((int64_t)L[j * n + idx] * x->data[j]) >> Q16_SHIFT;
        }
        x->data[idx] = q16_div((q16_t)sum, L[idx * n + idx]);
    }

    return SYN_OK;
}

SYN_Status syn_matrix_solve_cholesky(const SYN_Matrix *A, const SYN_Matrix *b, SYN_Matrix *x)
{
    SYN_ASSERT(A != NULL && b != NULL && x != NULL);
    uint8_t n = A->rows;
    if (n > SYN_SOLVER_MAX_N)
        return SYN_INVALID_PARAM;
    q16_t L[n * n];
    q16_t y[n];
    return syn_matrix_solve_cholesky_work(A, b, x, L, y);
}

SYN_Status syn_matrix_least_squares_work(const SYN_Matrix *A, const SYN_Matrix *b, SYN_Matrix *x,
                                         q16_t *ata_data, q16_t *atb_data, q16_t *at_data,
                                         q16_t *solver_lu, q16_t *solver_y)
{
    SYN_ASSERT(A != NULL && b != NULL && x != NULL);

    uint8_t m = A->rows;
    uint8_t n = A->cols;
    if (m < n || b->rows != m || b->cols != 1 || x->rows != n || x->cols != 1) {
        return SYN_INVALID_PARAM;
    }

    SYN_Matrix AtA = {ata_data, n, n};
    SYN_Matrix Atb = {atb_data, n, 1};
    SYN_Matrix AT = {at_data, n, m};
    memset(at_data, 0, (size_t)n * m * sizeof(q16_t));
    memset(ata_data, 0, (size_t)n * n * sizeof(q16_t));
    memset(atb_data, 0, (size_t)n * sizeof(q16_t));

    syn_matrix_transpose(A, &AT);
    syn_matrix_mul(&AT, A, &AtA);
    syn_matrix_mul(&AT, b, &Atb);

    SYN_Status status = syn_matrix_solve_cholesky_work(&AtA, &Atb, x, solver_lu, solver_y);
    if (status != SYN_OK) {
        uint8_t P_temp[SYN_SOLVER_MAX_N];
        status = syn_matrix_solve_lu_work(&AtA, &Atb, x, solver_lu, P_temp, solver_y);
    }

    return status;
}

SYN_Status syn_matrix_least_squares(const SYN_Matrix *A, const SYN_Matrix *b, SYN_Matrix *x)
{
    SYN_ASSERT(A != NULL && b != NULL && x != NULL);
    uint8_t m = A->rows;
    uint8_t n = A->cols;
    if (m < n)
        return SYN_INVALID_PARAM;

    q16_t ata_data[n * n];
    q16_t atb_data[n];
    q16_t at_data[n * m];
    q16_t solver_lu[n * n];
    q16_t solver_y[n];

    return syn_matrix_least_squares_work(A, b, x, ata_data, atb_data, at_data, solver_lu, solver_y);
}

SYN_Status syn_matrix_get_block(const SYN_Matrix *src, uint8_t r0, uint8_t c0, SYN_Matrix *dst)
{
    SYN_ASSERT(src != NULL && dst != NULL);
    if (r0 + dst->rows > src->rows || c0 + dst->cols > src->cols) {
        return SYN_INVALID_PARAM;
    }

    for (uint8_t r = 0; r < dst->rows; r++) {
        for (uint8_t c = 0; c < dst->cols; c++) {
            SYN_MAT_AT(dst, r, c) = SYN_MAT_AT(src, r0 + r, c0 + c);
        }
    }
    return SYN_OK;
}

SYN_Status syn_matrix_set_block(SYN_Matrix *dst, uint8_t r0, uint8_t c0, const SYN_Matrix *src)
{
    SYN_ASSERT(dst != NULL && src != NULL);
    if (r0 + src->rows > dst->rows || c0 + src->cols > dst->cols) {
        return SYN_INVALID_PARAM;
    }

    for (uint8_t r = 0; r < src->rows; r++) {
        for (uint8_t c = 0; c < src->cols; c++) {
            SYN_MAT_AT(dst, r0 + r, c0 + c) = SYN_MAT_AT(src, r, c);
        }
    }
    return SYN_OK;
}

SYN_Status syn_matrix_outer_product(const q16_t *u, uint8_t rows, const q16_t *v, uint8_t cols,
                                    SYN_Matrix *out)
{
    if (u == NULL || v == NULL || out == NULL)
        return SYN_INVALID_PARAM;
    if (out->rows != rows || out->cols != cols)
        return SYN_INVALID_PARAM;

    for (uint8_t r = 0; r < rows; r++) {
        for (uint8_t c = 0; c < cols; c++) {
            SYN_MAT_AT(out, r, c) = q16_mul(u[r], v[c]);
        }
    }
    return SYN_OK;
}

SYN_Status syn_matrix_qr(const SYN_Matrix *A, SYN_Matrix *Q, SYN_Matrix *R)
{
    SYN_ASSERT(A != NULL && Q != NULL && R != NULL);

    uint8_t m = A->rows;
    uint8_t n = A->cols;
    if (m < n || Q->rows != m || Q->cols != n || R->rows != n || R->cols != n) {
        return SYN_INVALID_PARAM;
    }

    syn_matrix_zero(R);
    syn_matrix_copy(Q, A);

    /* Modified Gram-Schmidt orthogonalization */
    q16_t v[SYN_SOLVER_MAX_N];
    if (m > SYN_SOLVER_MAX_N)
        return SYN_INVALID_PARAM;

    for (uint8_t k = 0; k < n; k++) {
        /* Extract k-th column into v */
        for (uint8_t i = 0; i < m; i++) {
            v[i] = SYN_MAT_AT(Q, i, k);
        }

        /* Orthogonalize against previous q_j columns */
        for (uint8_t j = 0; j < k; j++) {
            q16_t r_jk = 0;
            int64_t dot = 0;
            for (uint8_t i = 0; i < m; i++) {
                dot += (int64_t)SYN_MAT_AT(Q, i, j) * v[i];
            }
            r_jk = (q16_t)(dot >> Q16_SHIFT);
            SYN_MAT_AT(R, j, k) = r_jk;

            for (uint8_t i = 0; i < m; i++) {
                v[i] -= q16_mul(r_jk, SYN_MAT_AT(Q, i, j));
            }
        }

        /* Compute norm of v */
        q16_t norm_v = syn_vec_norm(v, m);
        if (norm_v == 0)
            return SYN_ERROR; /* Singular / linearly dependent */

        SYN_MAT_AT(R, k, k) = norm_v;

        /* Normalize and store in Q */
        for (uint8_t i = 0; i < m; i++) {
            SYN_MAT_AT(Q, i, k) = q16_div(v[i], norm_v);
        }
    }

    return SYN_OK;
}

SYN_Status syn_matrix_eigen_sym2(const SYN_Matrix *A, q16_t evals[2], SYN_Matrix *E)
{
    if (A == NULL || evals == NULL || E == NULL)
        return SYN_INVALID_PARAM;
    if (A->rows != 2 || A->cols != 2 || E->rows != 2 || E->cols != 2)
        return SYN_INVALID_PARAM;

    q16_t a = SYN_MAT_AT(A, 0, 0);
    q16_t b = SYN_MAT_AT(A, 0, 1);
    q16_t d = SYN_MAT_AT(A, 1, 1);

    q16_t trace = a + d;
    q16_t diff = a - d;

    q16_t disc = q16_hypot(diff, q16_mul(Q16_FROM_INT(2), b));
    q16_t l1 = (trace + disc) >> 1;
    q16_t l2 = (trace - disc) >> 1;

    evals[0] = l1;
    evals[1] = l2;

    /* Eigenvectors */
    if (b == 0) {
        SYN_MAT_AT(E, 0, 0) = Q16_ONE;
        SYN_MAT_AT(E, 0, 1) = 0;
        SYN_MAT_AT(E, 1, 0) = 0;
        SYN_MAT_AT(E, 1, 1) = Q16_ONE;
    } else {
        q16_t v1[2] = {l1 - d, b};
        q16_t v2[2] = {l2 - d, b};
        q16_t e1[2] = {0, 0};
        q16_t e2[2] = {0, 0};
        syn_vec_normalize(v1, e1, 2);
        syn_vec_normalize(v2, e2, 2);

        SYN_MAT_AT(E, 0, 0) = e1[0];
        SYN_MAT_AT(E, 0, 1) = e2[0];
        SYN_MAT_AT(E, 1, 0) = e1[1];
        SYN_MAT_AT(E, 1, 1) = e2[1];
    }

    return SYN_OK;
}

SYN_Status syn_matrix_eigen_sym3(const SYN_Matrix *A, q16_t evals[3], SYN_Matrix *E)
{
    if (A == NULL || evals == NULL || E == NULL)
        return SYN_INVALID_PARAM;
    if (A->rows != 3 || A->cols != 3 || E->rows != 3 || E->cols != 3)
        return SYN_INVALID_PARAM;

    SYN_MAT_DECL(S, 3, 3);
    syn_matrix_copy(&S, A);
    syn_matrix_identity(E);

    /* Cyclic Jacobi rotation algorithm for symmetric 3×3 matrix */
    for (uint8_t iter = 0; iter < 30; iter++) {
        /* Find largest off-diagonal element */
        uint8_t p = 0, q = 1;
        q16_t max_off = q16_abs(SYN_MAT_AT(&S, 0, 1));
        if (q16_abs(SYN_MAT_AT(&S, 0, 2)) > max_off) {
            max_off = q16_abs(SYN_MAT_AT(&S, 0, 2));
            p = 0;
            q = 2;
        }
        if (q16_abs(SYN_MAT_AT(&S, 1, 2)) > max_off) {
            max_off = q16_abs(SYN_MAT_AT(&S, 1, 2));
            p = 1;
            q = 2;
        }

        if (max_off < 4)
            break; /* Converged (less than 1e-4 in Q16) */

        q16_t app = SYN_MAT_AT(&S, p, p);
        q16_t aqq = SYN_MAT_AT(&S, q, q);
        q16_t apq = SYN_MAT_AT(&S, p, q);

        q16_t phi = q16_atan2(q16_mul(Q16_FROM_INT(2), apq), aqq - app) >> 1;
        q16_t c = q16_cos(phi);
        q16_t s = q16_sin(phi);

        /* Update matrix S = Jᵀ S J */
        for (uint8_t i = 0; i < 3; i++) {
            if (i != p && i != q) {
                q16_t a_ip = SYN_MAT_AT(&S, i, p);
                q16_t a_iq = SYN_MAT_AT(&S, i, q);
                SYN_MAT_AT(&S, i, p) = SYN_MAT_AT(&S, p, i) = q16_mul(c, a_ip) - q16_mul(s, a_iq);
                SYN_MAT_AT(&S, i, q) = SYN_MAT_AT(&S, q, i) = q16_mul(s, a_ip) + q16_mul(c, a_iq);
            }
        }
        SYN_MAT_AT(&S, p, p) = q16_mul(c, q16_mul(c, app) - q16_mul(s, apq)) -
                               q16_mul(s, q16_mul(c, apq) - q16_mul(s, aqq));
        SYN_MAT_AT(&S, q, q) = q16_mul(s, q16_mul(s, app) + q16_mul(c, apq)) +
                               q16_mul(c, q16_mul(s, apq) + q16_mul(c, aqq));
        SYN_MAT_AT(&S, p, q) = SYN_MAT_AT(&S, q, p) = 0;

        /* Update eigenvector matrix E = E J */
        for (uint8_t i = 0; i < 3; i++) {
            q16_t e_ip = SYN_MAT_AT(E, i, p);
            q16_t e_iq = SYN_MAT_AT(E, i, q);
            SYN_MAT_AT(E, i, p) = q16_mul(c, e_ip) - q16_mul(s, e_iq);
            SYN_MAT_AT(E, i, q) = q16_mul(s, e_ip) + q16_mul(c, e_iq);
        }
    }

    evals[0] = SYN_MAT_AT(&S, 0, 0);
    evals[1] = SYN_MAT_AT(&S, 1, 1);
    evals[2] = SYN_MAT_AT(&S, 2, 2);

    /* Sort eigenvalues and eigenvectors descending */
    for (uint8_t i = 0; i < 2; i++) {
        for (uint8_t j = i + 1; j < 3; j++) {
            if (evals[j] > evals[i]) {
                q16_t tmp_e = evals[i];
                evals[i] = evals[j];
                evals[j] = tmp_e;
                for (uint8_t k = 0; k < 3; k++) {
                    q16_t tmp_v = SYN_MAT_AT(E, k, i);
                    SYN_MAT_AT(E, k, i) = SYN_MAT_AT(E, k, j);
                    SYN_MAT_AT(E, k, j) = tmp_v;
                }
            }
        }
    }

    return SYN_OK;
}

#endif /* SYN_USE_MATRIX */