Skip to main content

HOLOGRAPHIC_BATCH_SIMILARITY

Constant HOLOGRAPHIC_BATCH_SIMILARITY 

Source
pub const HOLOGRAPHIC_BATCH_SIMILARITY: &str = r#"
struct TDC {
    tropical: f32,
    dual_real: f32,
    dual_dual: f32,
    clifford: array<f32, 8>,
    _padding: array<f32, 5>,
}

@group(0) @binding(0) var<storage, read> vectors_a: array<TDC>;
@group(0) @binding(1) var<storage, read> vectors_b: array<TDC>;
@group(0) @binding(2) var<storage, read_write> similarities: array<f32>;
@group(0) @binding(3) var<uniform> params: vec4<u32>; // [count_a, count_b, mode, 0]
                                                          // mode: 0=pairwise (a[i] vs b[i]), 1=matrix (all pairs)

// Compute reverse of multivector (flip sign of grades 2 and 3)
fn reverse_sign(grade: u32) -> f32 {
    switch grade {
        case 0u, 1u, 2u, 3u: { return 1.0; }
        default: { return -1.0; }
    }
}

// Compute scalar product <A B̃>₀ - the proper inner product for similarity
fn scalar_product_with_reverse(a: TDC, b: TDC) -> f32 {
    return a.clifford[0] * b.clifford[0] * reverse_sign(0u)
         + a.clifford[1] * b.clifford[1] * reverse_sign(1u)
         + a.clifford[2] * b.clifford[2] * reverse_sign(2u)
         + a.clifford[3] * b.clifford[3] * reverse_sign(3u)
         + a.clifford[4] * b.clifford[4] * reverse_sign(4u)
         + a.clifford[5] * b.clifford[5] * reverse_sign(5u)
         + a.clifford[6] * b.clifford[6] * reverse_sign(6u)
         + a.clifford[7] * b.clifford[7] * reverse_sign(7u);
}

fn norm(v: TDC) -> f32 {
    let sum = v.clifford[0] * v.clifford[0]
            + v.clifford[1] * v.clifford[1]
            + v.clifford[2] * v.clifford[2]
            + v.clifford[3] * v.clifford[3]
            + v.clifford[4] * v.clifford[4]
            + v.clifford[5] * v.clifford[5]
            + v.clifford[6] * v.clifford[6]
            + v.clifford[7] * v.clifford[7];
    return sqrt(sum);
}

@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
    let idx = global_id.x;
    let count_a = params[0];
    let count_b = params[1];
    let mode = params[2];

    if (mode == 0u) {
        // Pairwise mode: similarities[i] = sim(a[i], b[i])
        if (idx >= count_a) {
            return;
        }

        let a = vectors_a[idx];
        let b = vectors_b[idx];

        let norm_a = norm(a);
        let norm_b = norm(b);

        if (norm_a < 1e-10 || norm_b < 1e-10) {
            similarities[idx] = 0.0;
            return;
        }

        let inner = scalar_product_with_reverse(a, b);
        similarities[idx] = inner / (norm_a * norm_b);
    } else {
        // Matrix mode: similarities[i * count_b + j] = sim(a[i], b[j])
        let total = count_a * count_b;
        if (idx >= total) {
            return;
        }

        let i = idx / count_b;
        let j = idx % count_b;

        let a = vectors_a[i];
        let b = vectors_b[j];

        let norm_a = norm(a);
        let norm_b = norm(b);

        if (norm_a < 1e-10 || norm_b < 1e-10) {
            similarities[idx] = 0.0;
            return;
        }

        let inner = scalar_product_with_reverse(a, b);
        similarities[idx] = inner / (norm_a * norm_b);
    }
}
"#;
Expand description

Batch similarity computation for holographic vectors Computes pairwise similarities using inner product with reverse: <A B̃>₀