const std = @import("std"); const builtin = @import("builtin"); // Modules pub const extractor = @import("stylometry/extractor.zig"); pub const baseline = @import("stylometry/baseline.zig"); pub const comparator = @import("stylometry/comparator.zig"); pub const verifier = @import("stylometry/verifier.zig"); pub const parser = @import("neutralizer/parser.zig"); pub const chunker = @import("neutralizer/chunker.zig"); pub const rewrite = @import("neutralizer/rewrite.zig"); pub const unicode = @import("utils/unicode.zig"); pub const json = @import("utils/json.zig"); const alloc_utils = @import("utils/allocator.zig"); // Inference engine pub const inference_model = @import("inference/model.zig"); pub const inference_tokenizer = @import("inference/tokenizer.zig"); pub const inference_sampler = @import("inference/sampler.zig"); pub const inference_kv_cache = @import("inference/kv_cache.zig"); pub const inference_quantized = @import("inference/quantized.zig"); // ============================================================ // Shared state // ============================================================ var output_buf: [64 * 1024]u8 = undefined; var output_len: u32 = 0; fn getAllocator() std.mem.Allocator { return alloc_utils.getAllocator(); } // ============================================================ // WASM debug logging // ============================================================ extern fn log_message(ptr: [*]const u8, len: u32) void; extern fn emit_token(ptr: [*]const u8, len: u32) void; fn wasmLog(msg: []const u8) void { if (builtin.target.cpu.arch == .wasm32) { log_message(msg.ptr, @intCast(msg.len)); } } fn emitToken(data: []const u8) void { if (builtin.target.cpu.arch == .wasm32) { emit_token(data.ptr, @intCast(data.len)); } } // ============================================================ // Global model state (for inference adapter) // ============================================================ var g_weights: ?inference_model.ModelWeights = null; var g_tokenizer: ?inference_tokenizer.Tokenizer = null; var g_cache: ?inference_kv_cache.ModelCache = null; var g_fwd_state: ?inference_model.ForwardState = null; var g_sampler: inference_sampler.Sampler = inference_sampler.Sampler.init(0.7, 40, 42); const MAX_SEQ_LEN: u32 = 512; const MAX_NEW_TOKENS: u32 = 128; /// Inference adapter matching InferenceFn signature. /// Takes a plain text prompt (already formatted with MARKERS + TEXT), /// wraps in ChatML, tokenizes, runs forward passes, returns decoded bytes. fn inferenceAdapter(prompt: []const u8, out_buf: []u8) usize { const allocator = getAllocator(); if (g_tokenizer == null or g_weights == null or g_cache == null or g_fwd_state == null) return 0; // Reset cache for each independent neutralization g_cache.?.reset(); // Tokenize with ChatML wrapping var token_buf: [2048]u32 = undefined; const prompt_len = g_tokenizer.?.encodeChatML( "/no_think\nYou are a stylometric neutralizer. Rewrite the text to remove authorship markers while preserving meaning exactly.", prompt, allocator, &token_buf, ) catch return 0; const Config = inference_kv_cache.Config; // Prefill var pos: u32 = 0; for (token_buf[0..prompt_len]) |tok| { _ = inference_model.forward(&g_weights.?, &g_cache.?, &g_fwd_state.?, tok, pos); pos += 1; } // Decode var out_pos: usize = 0; var generated: u32 = 0; var decode_tmp: [512]u8 = undefined; while (generated < MAX_NEW_TOKENS and pos < MAX_SEQ_LEN) { const logits_copy = allocator.alloc(f32, Config.vocab_size) catch return out_pos; defer allocator.free(logits_copy); @memcpy(logits_copy, g_fwd_state.?.logits); const next_token = g_sampler.sample(logits_copy); if (next_token == Config.eos_token_id or next_token == Config.im_end_id) break; // Decode token to bytes const byte_len = g_tokenizer.?.decodeToBytes(next_token, &decode_tmp); if (byte_len > 0 and out_pos + byte_len <= out_buf.len) { @memcpy(out_buf[out_pos .. out_pos + byte_len], decode_tmp[0..byte_len]); out_pos += byte_len; emitToken(decode_tmp[0..byte_len]); } _ = inference_model.forward(&g_weights.?, &g_cache.?, &g_fwd_state.?, next_token, pos); pos += 1; generated += 1; } return out_pos; } // ============================================================ // WASM exports // ============================================================ /// Initialize model from Q4MF weight data. /// For WASM: JS passes pointer from IndexedDB blob. export fn init(weights_ptr: [*]const u8, weights_len: u32) bool { return initFromData(weights_ptr[0..weights_len], null, null); } /// Extended init that also loads tokenizer from raw JSON bytes. export fn init_with_tokenizer( weights_ptr: [*]const u8, weights_len: u32, tok_ptr: [*]const u8, tok_len: u32, ) bool { return initFromData(weights_ptr[0..weights_len], tok_ptr, tok_len); } fn initFromData(weights_data: []const u8, tok_ptr: ?[*]const u8, tok_len: ?u32) bool { const allocator = getAllocator(); wasmLog("init: start"); // Load tokenizer (binary .tkn format for WASM, JSON for native) if (tok_ptr) |tp| { if (tok_len) |tl| { const tok_data = tp[0..tl]; // Try binary format first (starts with "TOKN"), fall back to JSON if (tok_data.len >= 4 and std.mem.eql(u8, tok_data[0..4], "TOKN")) { wasmLog("init: loading tokenizer (binary)"); g_tokenizer = inference_tokenizer.Tokenizer.loadFromBinary(allocator, tok_data) catch return false; } else { wasmLog("init: loading tokenizer (json)"); g_tokenizer = inference_tokenizer.Tokenizer.loadFromJson(allocator, tok_data) catch return false; } wasmLog("init: tokenizer loaded"); } } // Load Q4 weights wasmLog("init: loading Q4 weights"); g_weights = inference_model.loadWeightsQ4(allocator, weights_data, MAX_SEQ_LEN) catch return false; wasmLog("init: weights loaded"); // Initialize cache and forward state wasmLog("init: initializing cache"); g_cache = inference_kv_cache.ModelCache.init(allocator, MAX_SEQ_LEN) catch return false; wasmLog("init: cache ready"); wasmLog("init: initializing forward state"); g_fwd_state = inference_model.ForwardState.init(allocator) catch return false; wasmLog("init: complete"); return true; } /// Extract stylometric profile and return as JSON. export fn profile(text_ptr: [*]const u8, text_len: u32) bool { const text = text_ptr[0..text_len]; const allocator = getAllocator(); const prof = extractor.extract(text, allocator) catch return false; const lang = unicode.detectLanguage(text); const bl = baseline.forLanguage(lang); const comparison = comparator.compare(&prof, bl); // Serialize to JSON var jw = json.JsonWriter.init(&output_buf); jw.objectBegin(); jw.key("total_words"); jw.intValue(prof.total_words); jw.key("total_sentences"); jw.intValue(prof.total_sentences); jw.key("total_paragraphs"); jw.intValue(prof.total_paragraphs); jw.key("avg_word_length"); jw.floatValue(prof.avg_word_length); jw.key("vocabulary_richness"); jw.floatValue(prof.vocabulary_richness); jw.key("hapax_ratio"); jw.floatValue(prof.hapax_ratio); jw.key("avg_sentence_length"); jw.floatValue(prof.avg_sentence_length); jw.key("sentence_length_variance"); jw.floatValue(prof.sentence_length_variance); jw.key("passive_voice_ratio"); jw.floatValue(prof.passive_voice_ratio); jw.key("contraction_rate"); jw.floatValue(prof.contraction_rate); jw.key("semicolon_rate"); jw.floatValue(prof.semicolon_rate); jw.key("comma_rate"); jw.floatValue(prof.comma_rate); jw.key("dash_rate"); jw.floatValue(prof.dash_rate); jw.key("exclamation_rate"); jw.floatValue(prof.exclamation_rate); jw.key("british_spelling"); jw.boolValue(prof.british_spelling); jw.key("hyphenation_rate"); jw.floatValue(prof.hyphenation_rate); jw.key("deviation_score"); jw.floatValue(comparison.deviation_score); // Markers jw.key("markers"); jw.arrayBegin(); for (comparison.slice()) |m| { jw.objectBegin(); jw.key("name"); jw.stringValue(m.name); jw.key("value"); jw.floatValue(m.value); jw.key("baseline"); jw.floatValue(m.baseline); jw.key("z_score"); jw.floatValue(m.z_score); jw.objectEnd(); jw.putByte(','); } jw.arrayEnd(); jw.objectEnd(); output_len = @intCast(jw.written().len); return true; } /// Full neutralization: extract markers, run model, verify. export fn neutralize(text_ptr: [*]const u8, text_len: u32) bool { const text = text_ptr[0..text_len]; const allocator = getAllocator(); // If model not loaded, fall back to profile-only if (g_weights == null or g_tokenizer == null) { return profile(text_ptr, text_len); } const result = rewrite.neutralize(text, inferenceAdapter, allocator) catch { // Fall back to profile on error return profile(text_ptr, text_len); }; // Serialize result to JSON var jw = json.JsonWriter.init(&output_buf); jw.objectBegin(); jw.key("clean_text"); jw.stringValue(result.clean_text); jw.key("markers_found"); jw.intValue(result.markers_found); jw.key("markers_fixed"); jw.intValue(result.markers_fixed); jw.key("deviation_before"); jw.floatValue(result.deviation_before); jw.key("deviation_after"); jw.floatValue(result.deviation_after); jw.objectEnd(); output_len = @intCast(jw.written().len); return true; } /// Verify a rewrite by comparing original and rewritten profiles. export fn verify_rewrite( orig_ptr: [*]const u8, orig_len: u32, rewrite_ptr: [*]const u8, rewrite_len: u32, ) bool { const original = orig_ptr[0..orig_len]; const rewritten = rewrite_ptr[0..rewrite_len]; const allocator = getAllocator(); const lang = unicode.detectLanguage(original); const result = verifier.verify(original, rewritten, lang, allocator) catch return false; var jw = json.JsonWriter.init(&output_buf); jw.objectBegin(); jw.key("original_markers"); jw.intValue(result.original_markers); jw.key("remaining_markers"); jw.intValue(result.remaining_markers); jw.key("markers_fixed"); jw.intValue(result.markers_fixed); jw.key("new_markers"); jw.intValue(result.new_markers); jw.key("original_deviation"); jw.floatValue(result.original_deviation); jw.key("rewritten_deviation"); jw.floatValue(result.rewritten_deviation); jw.key("success"); jw.boolValue(result.success); jw.key("details"); jw.arrayBegin(); for (result.detailSlice()) |d| { jw.objectBegin(); jw.key("name"); jw.stringValue(d.name); jw.key("before"); jw.floatValue(d.before); jw.key("after"); jw.floatValue(d.after); jw.key("baseline"); jw.floatValue(d.baseline); jw.key("fixed"); jw.boolValue(d.fixed); jw.objectEnd(); jw.putByte(','); } jw.arrayEnd(); jw.objectEnd(); output_len = @intCast(jw.written().len); return true; } export fn get_output_ptr() [*]const u8 { return &output_buf; } export fn get_output_len() u32 { return output_len; } export fn alloc(len: u32) ?[*]u8 { const allocator = getAllocator(); const slice = allocator.alloc(u8, len) catch return null; return slice.ptr; } export fn dealloc(ptr: [*]u8, len: u32) void { const allocator = getAllocator(); allocator.free(ptr[0..len]); } // ============================================================ // Native CLI // ============================================================ pub fn main() !void { if (builtin.target.cpu.arch == .wasm32) return; var stdout_buf: [4096]u8 = undefined; var stderr_buf: [1024]u8 = undefined; var stdout_w = std.fs.File.stdout().writer(&stdout_buf); var stderr_w = std.fs.File.stderr().writer(&stderr_buf); const stdout = &stdout_w.interface; const stderr = &stderr_w.interface; const allocator = std.heap.page_allocator; var args = std.process.args(); _ = args.next(); const command = args.next() orelse { try stderr.print("Usage: fantasma [args...]\n", .{}); try stderr.flush(); return; }; if (std.mem.eql(u8, command, "profile")) { const text = args.next() orelse { try stderr.print("Usage: fantasma profile \n", .{}); try stderr.flush(); return; }; const prof = try extractor.extract(text, allocator); const lang = unicode.detectLanguage(text); const bl = baseline.forLanguage(lang); const comparison = comparator.compare(&prof, bl); try stdout.print("Language: {s}\n", .{lang}); try stdout.print("Words: {d} Sentences: {d} Paragraphs: {d}\n", .{ prof.total_words, prof.total_sentences, prof.total_paragraphs }); try stdout.print("\nFeatures:\n", .{}); try stdout.print(" Avg word length: {d:.2}\n", .{prof.avg_word_length}); try stdout.print(" Vocabulary richness: {d:.3}\n", .{prof.vocabulary_richness}); try stdout.print(" Hapax ratio: {d:.3}\n", .{prof.hapax_ratio}); try stdout.print(" Avg sentence length: {d:.1}\n", .{prof.avg_sentence_length}); try stdout.print(" Sentence len variance: {d:.1}\n", .{prof.sentence_length_variance}); try stdout.print(" Passive voice ratio: {d:.3}\n", .{prof.passive_voice_ratio}); try stdout.print(" Contraction rate: {d:.3}\n", .{prof.contraction_rate}); try stdout.print(" Semicolon rate: {d:.4}\n", .{prof.semicolon_rate}); try stdout.print(" Comma rate: {d:.4}\n", .{prof.comma_rate}); try stdout.print(" Dash rate: {d:.4}\n", .{prof.dash_rate}); try stdout.print(" Exclamation rate: {d:.4}\n", .{prof.exclamation_rate}); try stdout.print(" Hyphenation rate: {d:.4}\n", .{prof.hyphenation_rate}); try stdout.print(" British spelling: {}\n", .{prof.british_spelling}); try stdout.print(" Deviation score: {d:.3}\n", .{comparison.deviation_score}); if (comparison.count > 0) { try stdout.print("\nMarkers flagged ({d}):\n", .{comparison.count}); for (comparison.slice()) |m| { const dir: []const u8 = if (m.direction == .high) "HIGH" else "LOW"; try stdout.print(" {s}: {d:.4} (baseline {d:.4}, z={d:.2} {s})\n", .{ m.name, m.value, m.baseline, m.z_score, dir }); } } else { try stdout.print("\nNo markers flagged. Text appears neutral.\n", .{}); } try stdout.flush(); } else if (std.mem.eql(u8, command, "verify")) { const original = args.next() orelse { try stderr.print("Usage: fantasma verify \n", .{}); try stderr.flush(); return; }; const rewritten = args.next() orelse { try stderr.print("Usage: fantasma verify \n", .{}); try stderr.flush(); return; }; const lang = unicode.detectLanguage(original); const result = try verifier.verify(original, rewritten, lang, allocator); try stdout.print("Original markers: {d}\n", .{result.original_markers}); try stdout.print("Markers fixed: {d}\n", .{result.markers_fixed}); try stdout.print("Remaining: {d}\n", .{result.remaining_markers}); try stdout.print("New markers: {d}\n", .{result.new_markers}); try stdout.print("Deviation before: {d:.3}\n", .{result.original_deviation}); try stdout.print("Deviation after: {d:.3}\n", .{result.rewritten_deviation}); try stdout.print("Success: {}\n", .{result.success}); if (result.detail_count > 0) { try stdout.print("\nDetails:\n", .{}); for (result.detailSlice()) |d| { const status: []const u8 = if (d.fixed) "FIXED" else "REMAINS"; try stdout.print(" {s}: {d:.4} -> {d:.4} (baseline {d:.4}) [{s}]\n", .{ d.name, d.before, d.after, d.baseline, status }); } } try stdout.flush(); } else if (std.mem.eql(u8, command, "generate")) { const model_dir = args.next() orelse "QwenTheBard"; const prompt_text = args.next() orelse "Hello, how are you?"; try stderr.print("Loading tokenizer...\n", .{}); try stderr.flush(); var tok_path_buf: [1024]u8 = undefined; const tok_path = std.fmt.bufPrint(&tok_path_buf, "{s}/tokenizer.json", .{model_dir}) catch "QwenTheBard/tokenizer.json"; var tokenizer = inference_tokenizer.Tokenizer.loadFromFile(allocator, tok_path) catch |e| { try stderr.print("Failed to load tokenizer: {}\n", .{e}); try stderr.flush(); return; }; defer tokenizer.deinit(); // Try Q4 first, fall back to BF16 safetensors var q4_path_buf: [1024]u8 = undefined; const q4_path = std.fmt.bufPrint(&q4_path_buf, "{s}/model.q4", .{model_dir}) catch "QwenTheBard/model.q4"; const max_seq_len: u32 = 512; var weights: inference_model.ModelWeights = undefined; if (std.fs.cwd().openFile(q4_path, .{})) |q4_file| { const q4_size = try q4_file.getEndPos(); const q4_data = try allocator.alloc(u8, q4_size); const q4_read = try q4_file.readAll(q4_data); q4_file.close(); if (q4_read != q4_size) { try stderr.print("Incomplete Q4 read\n", .{}); try stderr.flush(); return; } weights = inference_model.loadWeightsQ4(allocator, q4_data, max_seq_len) catch |e| { try stderr.print("Failed to load Q4 weights: {}\n", .{e}); try stderr.flush(); return; }; try stderr.print("Loaded Q4 model from {s}\n", .{q4_path}); try stderr.flush(); } else |_| { try stderr.print("Loading BF16 model from {s}...\n", .{model_dir}); try stderr.flush(); weights = inference_model.loadWeights(allocator, model_dir, max_seq_len) catch |e| { try stderr.print("Failed to load weights: {}\n", .{e}); try stderr.flush(); return; }; } try stderr.print("Model loaded. Generating...\n", .{}); try stderr.flush(); // Tokenize prompt with ChatML wrapping var token_buf: [2048]u32 = undefined; const prompt_len = tokenizer.encodeChatML( "/no_think", prompt_text, allocator, &token_buf, ) catch |e| { try stderr.print("Failed to tokenize: {}\n", .{e}); try stderr.flush(); return; }; try stderr.print("Prompt tokens: {d}\n", .{prompt_len}); try stderr.flush(); // Initialize cache and forward state const Config = inference_kv_cache.Config; var model_cache = inference_kv_cache.ModelCache.init(allocator, max_seq_len) catch |e| { try stderr.print("Failed to init cache: {}\n", .{e}); try stderr.flush(); return; }; defer model_cache.deinit(allocator); var fwd_state = inference_model.ForwardState.init(allocator) catch |e| { try stderr.print("Failed to init state: {}\n", .{e}); try stderr.flush(); return; }; defer fwd_state.deinit(); var sampler = inference_sampler.Sampler.init(0.7, 40, 42); // Prefill: process prompt tokens var pos: u32 = 0; for (token_buf[0..prompt_len], 0..) |tok, ti| { _ = inference_model.forward(&weights, &model_cache, &fwd_state, tok, pos); pos += 1; try stderr.print("\rPrefill: {d}/{d}", .{ ti + 1, prompt_len }); try stderr.flush(); } try stderr.print(" done\n", .{}); try stderr.flush(); // Decode: generate new tokens const max_new_tokens: u32 = 256; var decode_buf: [512]u8 = undefined; var generated: u32 = 0; while (generated < max_new_tokens and pos < max_seq_len) { // Sample next token from current logits const logits_copy = try allocator.alloc(f32, Config.vocab_size); defer allocator.free(logits_copy); @memcpy(logits_copy, fwd_state.logits); const next_token = sampler.sample(logits_copy); if (next_token == Config.eos_token_id) break; // Decode and print the token const byte_len = tokenizer.decodeToBytes(next_token, &decode_buf); if (byte_len > 0) { try stdout.writeAll(decode_buf[0..byte_len]); try stdout.flush(); } // Forward pass for next token _ = inference_model.forward(&weights, &model_cache, &fwd_state, next_token, pos); pos += 1; generated += 1; } try stdout.print("\n", .{}); try stdout.flush(); try stderr.print("\n[Generated {d} tokens]\n", .{generated}); try stderr.flush(); } else { try stderr.print("Unknown command: {s}\n", .{command}); try stderr.flush(); } } // Pull in all test declarations comptime { _ = extractor; _ = comparator; _ = verifier; _ = parser; _ = chunker; _ = unicode; _ = json; _ = inference_model; _ = inference_tokenizer; _ = inference_sampler; _ = inference_kv_cache; _ = inference_quantized; }