/*
 reinhard@finalmedia.de
 Do 9. Jul 22:24:46 CEST 2026
 Public Domain

 splat2fsplat.c
 liest splat files von stdin und schreibt fsplat files auf stdout

 musl-gcc -O3 -march=native -static-pie -fPIE splat2fsplat.c -o splat2fsplat -lm

*/

#include <math.h>
#include <stdint.h>
#include <unistd.h>

#define MAX_SPLATS 8500000
#define MAX_PATTERNS 256
#define KMEANS_ITERATIONS 15
#define SH_C0 0.28209479177387814f

typedef struct {
    float x, y, z;
    float sx, sy, sz;
    unsigned char r, g, b, opacity;
    unsigned char q0, q1, q2, q3;
} __attribute__((packed)) SFileIn;

typedef struct {
    float x, y, z, sx, sy, sz;
    unsigned char opacity, q0, q1, q2, q3;
    unsigned char sh_idx;
} __attribute__((packed)) FSplat;

SFileIn input_splats[MAX_SPLATS];
float rgb_data[MAX_SPLATS * 3];
uint8_t labels[MAX_SPLATS];

float rgb_centroids[MAX_PATTERNS * 3];
float new_centroids[MAX_PATTERNS * 3];
uint32_t counts[MAX_PATTERNS];
float sh_codebook[MAX_PATTERNS * 48];

#define OUT_BUF_COUNT 8192
FSplat out_buf[OUT_BUF_COUNT];

static inline void byte_zero(void* dst, size_t n) {
    char* d = (char*)dst;
    while (n--) *d++ = 0;
}

void run_kmeans_rgb(uint32_t n_splats) {
    for (int i = 0; i < MAX_PATTERNS; i++) {
        uint32_t idx = (i * (n_splats / MAX_PATTERNS)) % n_splats;
        rgb_centroids[i * 3 + 0] = rgb_data[idx * 3 + 0];
        rgb_centroids[i * 3 + 1] = rgb_data[idx * 3 + 1];
        rgb_centroids[i * 3 + 2] = rgb_data[idx * 3 + 2];
    }

    for (int iter = 0; iter < KMEANS_ITERATIONS; iter++) {
        byte_zero(new_centroids, MAX_PATTERNS * 3 * sizeof(float));
        byte_zero(counts, MAX_PATTERNS * sizeof(uint32_t));

        for (uint32_t i = 0; i < n_splats; i++) {
            float min_dist = 1e30f;
            uint8_t best_cluster = 0;
            float r = rgb_data[i * 3 + 0];
            float g = rgb_data[i * 3 + 1];
            float b = rgb_data[i * 3 + 2];

            for (int k = 0; k < MAX_PATTERNS; k++) {
                float dr = r - rgb_centroids[k * 3 + 0];
                float dg = g - rgb_centroids[k * 3 + 1];
                float db = b - rgb_centroids[k * 3 + 2];
                float dist = dr*dr + dg*dg + db*db;
                if (dist < min_dist) {
                    min_dist = dist;
                    best_cluster = k;
                }
            }
            labels[i] = best_cluster;
            counts[best_cluster]++;
            new_centroids[best_cluster * 3 + 0] += r;
            new_centroids[best_cluster * 3 + 1] += g;
            new_centroids[best_cluster * 3 + 2] += b;
        }

        for (int k = 0; k < MAX_PATTERNS; k++) {
            if (counts[k] > 0) {
                rgb_centroids[k * 3 + 0] = new_centroids[k * 3 + 0] / counts[k];
                rgb_centroids[k * 3 + 1] = new_centroids[k * 3 + 1] / counts[k];
                rgb_centroids[k * 3 + 2] = new_centroids[k * 3 + 2] / counts[k];
            }
        }
    }
}

int main(void) {
    uint32_t n_splats = 0;

    while (n_splats < MAX_SPLATS) {
        size_t bytes_read = read(0, &input_splats[n_splats], sizeof(SFileIn));
        if (bytes_read < sizeof(SFileIn)) break;
        n_splats++;
    }

    if (n_splats == 0) return 1;

    for (uint32_t i = 0; i < n_splats; i++) {
        rgb_data[i * 3 + 0] = (float)input_splats[i].r / 255.0f;
        rgb_data[i * 3 + 1] = (float)input_splats[i].g / 255.0f;
        rgb_data[i * 3 + 2] = (float)input_splats[i].b / 255.0f;
    }

    run_kmeans_rgb(n_splats);

    byte_zero(sh_codebook, MAX_PATTERNS * 48 * sizeof(float));
    for (int k = 0; k < MAX_PATTERNS; k++) {
        sh_codebook[k * 48 + 0]  = (rgb_centroids[k * 3 + 0] - 0.5f) / SH_C0;
        sh_codebook[k * 48 + 16] = (rgb_centroids[k * 3 + 1] - 0.5f) / SH_C0;
        sh_codebook[k * 48 + 32] = (rgb_centroids[k * 3 + 2] - 0.5f) / SH_C0;
    }

    uint32_t num_patterns = MAX_PATTERNS;
    write(1, &num_patterns, sizeof(uint32_t));
    write(1, sh_codebook, sizeof(float) * 48 * MAX_PATTERNS);
    write(1, &n_splats, sizeof(uint32_t));

    int out_p = 0;
    for (uint32_t i = 0; i < n_splats; i++) {
        out_buf[out_p].x = input_splats[i].x;
        out_buf[out_p].y = input_splats[i].y;
        out_buf[out_p].z = input_splats[i].z;
        out_buf[out_p].sx = input_splats[i].sx;
        out_buf[out_p].sy = input_splats[i].sy;
        out_buf[out_p].sz = input_splats[i].sz;
        out_buf[out_p].opacity = input_splats[i].opacity;
        out_buf[out_p].q0 = input_splats[i].q0;
        out_buf[out_p].q1 = input_splats[i].q1;
        out_buf[out_p].q2 = input_splats[i].q2;
        out_buf[out_p].q3 = input_splats[i].q3;
        out_buf[out_p].sh_idx = labels[i];

        out_p++;
        if (out_p >= OUT_BUF_COUNT) {
            write(1, out_buf, OUT_BUF_COUNT * sizeof(FSplat));
            out_p = 0;
        }
    }
    if (out_p > 0) write(1, out_buf, out_p * sizeof(FSplat));

    return 0;
}


