#include "../../src/TransformerRanker/TransformerRanker.h" #include #include #include #include #include #include #include #include #include #include #include namespace { namespace transformer = szilassi::transformer; namespace training = szilassi::training; constexpr std::size_t kHeaderBytes = 160; constexpr std::size_t kProbeInputHeaderBytes = 28; constexpr std::size_t kProbeOutputHeaderBytes = 20; constexpr std::size_t kProbeOutputFloatCount = 45; constexpr std::size_t kMaximumProbeCases = 64; constexpr std::uint32_t kProbeFormatVersion = 1; constexpr std::array kProbeInputMagic{{ 'S', 'Z', 'T', 'P', 'R', 'B', '0', '1'}}; constexpr std::array kProbeOutputMagic{{ 'S', 'Z', 'T', 'P', 'O', 'U', '0', '1'}}; int fail(int code, const std::string& message); std::uint32_t crc32(const std::uint8_t* data, std::size_t size) { std::uint32_t crc = 0xffffffffu; for (std::size_t index = 0; index < size; ++index) { crc ^= data[index]; for (int bit = 0; bit < 8; ++bit) { const std::uint32_t mask = 0u - (crc & 1u); crc = (crc >> 1) ^ (0xedb88320u & mask); } } return ~crc; } void put_u32(std::vector& bytes, std::size_t offset, std::uint32_t value) { for (int shift = 0; shift < 32; shift += 8) { bytes[offset++] = static_cast(value >> shift); } } void put_u64(std::vector& bytes, std::size_t offset, std::uint64_t value) { for (int shift = 0; shift < 64; shift += 8) { bytes[offset++] = static_cast(value >> shift); } } void put_float(std::vector& bytes, std::size_t offset, float value) { std::uint32_t bits = 0; std::memcpy(&bits, &value, sizeof(value)); put_u32(bytes, offset, bits); } std::vector make_zero_model() { std::vector bytes( static_cast(transformer::maximum_model_bytes()), 0); const std::array magic{{ 'S', 'Z', 'T', 'R', 'N', 'K', '0', '1'}}; std::copy(magic.begin(), magic.end(), bytes.begin()); put_u32(bytes, 8, transformer::kModelFormatVersion); put_u32(bytes, 12, transformer::kFeatureFormatVersion); put_u32(bytes, 16, transformer::kRequiredObjectiveVersion); put_u32(bytes, 20, 1); const std::array dimensions{{ static_cast(transformer::kFaceCount), static_cast(transformer::kFaceFeatureCount), static_cast(transformer::kGlobalFeatureCount), static_cast(transformer::kModelWidth), static_cast(transformer::kAttentionHeadCount), static_cast(transformer::kLayerCount), static_cast(transformer::kFeedForwardWidth), static_cast(transformer::kEnsembleSize), static_cast(transformer::kTopologyCount)}}; for (std::size_t index = 0; index < dimensions.size(); ++index) { put_u32(bytes, 24 + 4 * index, dimensions[index]); } put_u64(bytes, 64, 123456789ULL); bytes[72] = 0x42; put_float(bytes, 104, 0.5f); put_float(bytes, 108, 0.2f); put_float(bytes, 112, 0.1f); put_float(bytes, 116, 0.3f); const std::uint64_t payload_bytes = transformer::model_payload_float_count() * sizeof(float); put_u64(bytes, 120, payload_bytes); put_u32(bytes, 128, crc32(bytes.data() + kHeaderBytes, bytes.size() - kHeaderBytes)); put_u64(bytes, 136, transformer::model_payload_float_count()); put_u32(bytes, 132, 0); put_u32(bytes, 132, crc32(bytes.data(), kHeaderBytes)); return bytes; } bool write_file(const std::filesystem::path& path, const std::vector& bytes) { std::ofstream output(path, std::ios::binary | std::ios::trunc); return output && output.write( reinterpret_cast(bytes.data()), static_cast(bytes.size())); } bool read_bounded_file( const std::filesystem::path& path, std::size_t maximum_bytes, std::vector& bytes, std::string& error) { std::error_code size_error; const std::uintmax_t size = std::filesystem::file_size(path, size_error); if (size_error) { error = "cannot determine probe input size: " + size_error.message(); return false; } if (size > maximum_bytes || size > static_cast(std::numeric_limits::max())) { error = "probe input exceeds its fixed size bound"; return false; } bytes.resize(static_cast(size)); std::ifstream input(path, std::ios::binary); if (!input || !input.read( reinterpret_cast(bytes.data()), static_cast(bytes.size()))) { error = "cannot read probe input"; return false; } return true; } bool get_u32( const std::vector& bytes, std::size_t offset, std::uint32_t& value) { if (offset > bytes.size() || bytes.size() - offset < 4) return false; value = static_cast(bytes[offset]) | (static_cast(bytes[offset + 1]) << 8) | (static_cast(bytes[offset + 2]) << 16) | (static_cast(bytes[offset + 3]) << 24); return true; } bool get_float( const std::vector& bytes, std::size_t offset, float& value) { std::uint32_t bits = 0; if (!get_u32(bytes, offset, bits)) return false; std::memcpy(&value, &bits, sizeof(value)); return std::isfinite(value); } int run_parity_probe( const std::filesystem::path& model_path, const std::filesystem::path& input_path, const std::filesystem::path& output_path) { constexpr std::size_t feature_float_count = transformer::kFaceCount * transformer::kFaceFeatureCount + transformer::kGlobalFeatureCount; constexpr std::size_t input_record_bytes = sizeof(std::uint32_t) + feature_float_count * sizeof(float); constexpr std::size_t maximum_input_bytes = kProbeInputHeaderBytes + kMaximumProbeCases * input_record_bytes; std::vector input_bytes; std::string error; if (!read_bounded_file( input_path, maximum_input_bytes, input_bytes, error)) { return fail(20, error); } if (input_bytes.size() < kProbeInputHeaderBytes || !std::equal( kProbeInputMagic.begin(), kProbeInputMagic.end(), input_bytes.begin())) { return fail(21, "probe input magic or header is invalid"); } std::uint32_t version = 0; std::uint32_t case_count = 0; std::uint32_t face_count = 0; std::uint32_t face_feature_count = 0; std::uint32_t global_feature_count = 0; if (!get_u32(input_bytes, 8, version) || !get_u32(input_bytes, 12, case_count) || !get_u32(input_bytes, 16, face_count) || !get_u32(input_bytes, 20, face_feature_count) || !get_u32(input_bytes, 24, global_feature_count) || version != kProbeFormatVersion || case_count == 0 || case_count > kMaximumProbeCases || face_count != transformer::kFaceCount || face_feature_count != transformer::kFaceFeatureCount || global_feature_count != transformer::kGlobalFeatureCount) { return fail(22, "probe input dimensions or version are invalid"); } const std::size_t expected_input_bytes = kProbeInputHeaderBytes + static_cast(case_count) * input_record_bytes; if (input_bytes.size() != expected_input_bytes) { return fail(23, "probe input has trailing or truncated data"); } std::vector feature_cases(case_count); std::size_t input_offset = kProbeInputHeaderBytes; for (transformer::Features& features : feature_cases) { if (!get_u32(input_bytes, input_offset, features.topology) || features.topology >= transformer::kTopologyCount) { return fail(24, "probe topology is outside 0..58"); } input_offset += sizeof(std::uint32_t); for (auto& face : features.face) { for (float& value : face) { if (!get_float(input_bytes, input_offset, value)) { return fail(25, "probe face feature is truncated or non-finite"); } input_offset += sizeof(float); } } for (float& value : features.global) { if (!get_float(input_bytes, input_offset, value)) { return fail(26, "probe global feature is truncated or non-finite"); } input_offset += sizeof(float); } } transformer::TransformerRanker ranker; if (!ranker.load(model_path, &error)) return fail(27, error); const std::vector predictions = ranker.predict_batch(feature_cases); if (predictions.size() != feature_cases.size()) { return fail(28, "native probe returned an invalid batch size"); } constexpr std::size_t output_record_bytes = sizeof(std::uint32_t) + kProbeOutputFloatCount * sizeof(float); std::vector output_bytes( kProbeOutputHeaderBytes + predictions.size() * output_record_bytes, 0); std::copy(kProbeOutputMagic.begin(), kProbeOutputMagic.end(), output_bytes.begin()); put_u32(output_bytes, 8, kProbeFormatVersion); put_u32(output_bytes, 12, case_count); put_u32( output_bytes, 16, static_cast(kProbeOutputFloatCount)); std::size_t output_offset = kProbeOutputHeaderBytes; for (const transformer::Prediction& prediction : predictions) { put_u32(output_bytes, output_offset, prediction.finite ? 1u : 0u); output_offset += sizeof(std::uint32_t); const auto emit = [&](float value) { put_float(output_bytes, output_offset, value); output_offset += sizeof(float); }; emit(prediction.improvement_logit); emit(prediction.improvement_probability); emit(prediction.expected_defect_gain); emit(prediction.improvement_probability_variance); emit(prediction.expected_defect_gain_variance); for (float value : prediction.plane_probabilities) emit(value); for (float value : prediction.plane_probability_variance) emit(value); for (float value : prediction.move_probabilities) emit(value); for (float value : prediction.move_probability_variance) emit(value); for (float value : prediction.scale_probabilities) emit(value); for (float value : prediction.scale_probability_variance) emit(value); } if (!write_file(output_path, output_bytes)) { return fail(29, "cannot write native probe output"); } std::cout << "Transformer parity probe wrote " << predictions.size() << " bounded cases\n"; return 0; } bool approximately(float left, float right, float tolerance = 1.0e-5f) { return std::isfinite(left) && std::isfinite(right) && std::abs(left - right) <= tolerance; } int fail(int code, const std::string& message = {}) { std::cerr << "TransformerRankerSelfTest failure " << code; if (!message.empty()) std::cerr << ": " << message; std::cerr << '\n'; return code; } } // namespace int main(int argc, char** argv) { if (argc == 5 && std::string(argv[1]) == "--parity-probe") { return run_parity_probe(argv[2], argv[3], argv[4]); } if (argc != 1) { std::cerr << "Usage: TransformerRankerSelfTest " "[--parity-probe MODEL INPUT OUTPUT]\n"; return 64; } if (transformer::maximum_model_bytes() != kHeaderBytes + transformer::model_payload_float_count() * sizeof(float)) { return fail(1, "model byte accounting mismatch"); } transformer::FeatureInput input; input.topology = 42; input.guided_context = true; input.budget = {64, 8, 8, 128}; for (std::size_t face = 0; face < transformer::kFaceCount; ++face) { input.state.values[3 * face] = 1.0f + static_cast(face) * 0.01f; input.state.values[3 * face + 1] = 0.1f; input.state.values[3 * face + 2] = -0.05f; } input.metrics.crossings = 2; input.metrics.intersections = 3; input.metrics.crossing_loss = 0.25; input.metrics.intersection_loss = 0.5; input.metrics.geometry_penalty = 0.01; input.metrics.degeneracy_penalty = 0.02; input.metrics.worst_degeneracy = 0.03; input.metrics.energy = 5.0; input.metrics.min_plane_determinant = 1.0e-4; input.metrics.relative_min_edge = 0.1; input.metrics.min_turn_sine = 0.2; input.metrics.max_vertex_norm = 3.0; input.metrics.condition_number = 100.0; input.metrics.condition_penalty = 0.04; input.metrics.crossing_face_mask = 1u << 2; input.metrics.intersection_face_mask = 1u << 3; input.metrics.degeneracy_face_mask = 1u << 4; input.metrics.condition_face_mask = 1u << 5; input.metrics.flags = training::MetricPrecise | training::MetricDdVerified | training::MetricFinite | training::MetricHasSmoothLosses | training::MetricHasGeometryPenalty | training::MetricHasWorstDegeneracy | training::MetricHasShapeDescriptors | training::MetricHasCondition | training::MetricHasFaceMasks | training::MetricHasEnergy; transformer::Features features; std::string error; if (!transformer::make_features(input, features, &error) || features.topology != 42 || !approximately(features.face[2][8], 1.0f) || !approximately(features.face[3][9], 1.0f) || !approximately(features.global[0], 0.125f) || !approximately(features.global[1], 3.0f / 32.0f) || !approximately(features.global[22], 1.0f)) { return fail(2, error); } transformer::FeatureInput invalid = input; invalid.state.values[0] = std::numeric_limits::quiet_NaN(); if (transformer::make_features(invalid, features, &error)) { return fail(3, "non-finite plane was accepted"); } const std::filesystem::path directory = std::filesystem::path("runtime") / "transformer_ranker_selftest"; std::error_code directory_error; std::filesystem::create_directories(directory, directory_error); if (directory_error) return fail(4, directory_error.message()); const std::filesystem::path valid_path = directory / "zero_model.sztf"; const std::filesystem::path corrupt_path = directory / "corrupt_model.sztf"; std::vector model_bytes = make_zero_model(); if (!write_file(valid_path, model_bytes)) return fail(5, "cannot write valid model"); transformer::TransformerRanker ranker; if (!ranker.load(valid_path, &error) || !ranker.ready()) return fail(6, error); const transformer::Metadata metadata = ranker.metadata(); if (!metadata.approved || metadata.training_seed != 123456789ULL || metadata.training_digest[0] != 0x42 || !approximately(metadata.validation.average_precision, 0.5f)) { return fail(7, "metadata mismatch"); } const transformer::Prediction prediction = ranker.predict(features); if (!prediction.finite || !approximately(prediction.improvement_probability, 0.5f) || !approximately(prediction.expected_defect_gain, 0.0f) || !approximately(prediction.improvement_probability_variance, 0.0f)) { return fail(8, "zero model value output mismatch"); } for (float probability : prediction.plane_probabilities) { if (!approximately(probability, 1.0f / transformer::kFaceCount)) { return fail(9, "plane softmax mismatch"); } } const std::vector batch = ranker.predict_batch({features, features}); if (batch.size() != 2 || !batch[0].finite || !batch[1].finite) { return fail(10, "batch prediction failed"); } model_bytes.back() ^= 0x01; if (!write_file(corrupt_path, model_bytes)) return fail(11, "cannot write corrupt model"); if (ranker.load(corrupt_path, &error) || !ranker.ready() || !ranker.predict(features).finite) { return fail(12, "transactional corrupt-load fallback failed"); } if (ranker.load(directory / "missing.sztf", &error) || !ranker.ready()) { return fail(13, "transactional missing-load fallback failed"); } std::filesystem::remove(valid_path, directory_error); std::filesystem::remove(corrupt_path, directory_error); std::cout << "TransformerRankerSelfTest passed; model bytes=" << transformer::maximum_model_bytes() << '\n'; return 0; }