Transformer

This commit is contained in:
Efim Beshmenev
2026-07-12 21:48:44 +03:00
parent ffd46f89e3
commit 9e3dd7ce8b
11614 changed files with 16818 additions and 4458 deletions
@@ -0,0 +1,399 @@
#include "../../src/TransformerRanker/TransformerRanker.h"
#include <algorithm>
#include <array>
#include <cmath>
#include <cstdint>
#include <cstring>
#include <filesystem>
#include <fstream>
#include <iostream>
#include <limits>
#include <string>
#include <vector>
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<std::uint8_t, 8> kProbeInputMagic{{
'S', 'Z', 'T', 'P', 'R', 'B', '0', '1'}};
constexpr std::array<std::uint8_t, 8> 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<std::uint8_t>& bytes, std::size_t offset, std::uint32_t value) {
for (int shift = 0; shift < 32; shift += 8) {
bytes[offset++] = static_cast<std::uint8_t>(value >> shift);
}
}
void put_u64(std::vector<std::uint8_t>& bytes, std::size_t offset, std::uint64_t value) {
for (int shift = 0; shift < 64; shift += 8) {
bytes[offset++] = static_cast<std::uint8_t>(value >> shift);
}
}
void put_float(std::vector<std::uint8_t>& 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<std::uint8_t> make_zero_model() {
std::vector<std::uint8_t> bytes(
static_cast<std::size_t>(transformer::maximum_model_bytes()), 0);
const std::array<std::uint8_t, 8> 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<std::uint32_t, 9> dimensions{{
static_cast<std::uint32_t>(transformer::kFaceCount),
static_cast<std::uint32_t>(transformer::kFaceFeatureCount),
static_cast<std::uint32_t>(transformer::kGlobalFeatureCount),
static_cast<std::uint32_t>(transformer::kModelWidth),
static_cast<std::uint32_t>(transformer::kAttentionHeadCount),
static_cast<std::uint32_t>(transformer::kLayerCount),
static_cast<std::uint32_t>(transformer::kFeedForwardWidth),
static_cast<std::uint32_t>(transformer::kEnsembleSize),
static_cast<std::uint32_t>(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<std::uint8_t>& bytes) {
std::ofstream output(path, std::ios::binary | std::ios::trunc);
return output && output.write(
reinterpret_cast<const char*>(bytes.data()),
static_cast<std::streamsize>(bytes.size()));
}
bool read_bounded_file(
const std::filesystem::path& path,
std::size_t maximum_bytes,
std::vector<std::uint8_t>& 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::uintmax_t>(std::numeric_limits<std::size_t>::max())) {
error = "probe input exceeds its fixed size bound";
return false;
}
bytes.resize(static_cast<std::size_t>(size));
std::ifstream input(path, std::ios::binary);
if (!input || !input.read(
reinterpret_cast<char*>(bytes.data()),
static_cast<std::streamsize>(bytes.size()))) {
error = "cannot read probe input";
return false;
}
return true;
}
bool get_u32(
const std::vector<std::uint8_t>& bytes,
std::size_t offset,
std::uint32_t& value) {
if (offset > bytes.size() || bytes.size() - offset < 4) return false;
value = static_cast<std::uint32_t>(bytes[offset]) |
(static_cast<std::uint32_t>(bytes[offset + 1]) << 8) |
(static_cast<std::uint32_t>(bytes[offset + 2]) << 16) |
(static_cast<std::uint32_t>(bytes[offset + 3]) << 24);
return true;
}
bool get_float(
const std::vector<std::uint8_t>& 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<std::uint8_t> 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<std::size_t>(case_count) * input_record_bytes;
if (input_bytes.size() != expected_input_bytes) {
return fail(23, "probe input has trailing or truncated data");
}
std::vector<transformer::Features> 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<transformer::Prediction> 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<std::uint8_t> 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<std::uint32_t>(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<float>(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<float>::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<std::uint8_t> 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<transformer::Prediction> 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;
}
@@ -0,0 +1,65 @@
<?xml version="1.0" encoding="utf-8"?>
<Project DefaultTargets="Build" xmlns="http://schemas.microsoft.com/developer/msbuild/2003">
<ItemGroup Label="ProjectConfigurations">
<ProjectConfiguration Include="Release|x64">
<Configuration>Release</Configuration>
<Platform>x64</Platform>
</ProjectConfiguration>
<ProjectConfiguration Include="Debug|x64">
<Configuration>Debug</Configuration>
<Platform>x64</Platform>
</ProjectConfiguration>
</ItemGroup>
<PropertyGroup Label="Globals">
<VCProjectVersion>18.0</VCProjectVersion>
<Keyword>Win32Proj</Keyword>
<ProjectGuid>{FB8E7C11-3B60-47C2-962A-641913C1D38A}</ProjectGuid>
<RootNamespace>TransformerRankerSelfTest</RootNamespace>
<WindowsTargetPlatformVersion>10.0.26100.0</WindowsTargetPlatformVersion>
</PropertyGroup>
<Import Project="$(VCTargetsPath)\Microsoft.Cpp.Default.props" />
<PropertyGroup Condition="'$(Configuration)|$(Platform)'=='Debug|x64'" Label="Configuration">
<ConfigurationType>Application</ConfigurationType>
<UseDebugLibraries>true</UseDebugLibraries>
<PlatformToolset>v145</PlatformToolset>
<CharacterSet>Unicode</CharacterSet>
</PropertyGroup>
<PropertyGroup Condition="'$(Configuration)|$(Platform)'=='Release|x64'" Label="Configuration">
<ConfigurationType>Application</ConfigurationType>
<UseDebugLibraries>false</UseDebugLibraries>
<PlatformToolset>v145</PlatformToolset>
<WholeProgramOptimization>true</WholeProgramOptimization>
<CharacterSet>Unicode</CharacterSet>
</PropertyGroup>
<Import Project="$(VCTargetsPath)\Microsoft.Cpp.props" />
<PropertyGroup>
<RepositoryRoot>$([System.IO.Path]::GetFullPath('$(MSBuildThisFileDirectory)..\..\'))</RepositoryRoot>
<OutDir>$(RepositoryRoot)build\msbuild\bin\$(Platform)\$(Configuration)\</OutDir>
<IntDir>$(RepositoryRoot)build\msbuild\obj\$(ProjectName)\$(Platform)\$(Configuration)\</IntDir>
<LocalDebuggerWorkingDirectory>$(RepositoryRoot)</LocalDebuggerWorkingDirectory>
</PropertyGroup>
<ItemDefinitionGroup>
<ClCompile>
<WarningLevel>Level3</WarningLevel>
<SDLCheck>false</SDLCheck>
<ConformanceMode>true</ConformanceMode>
<LanguageStandard>stdcpp17</LanguageStandard>
<AdditionalOptions>/utf-8 %(AdditionalOptions)</AdditionalOptions>
<AdditionalIncludeDirectories>$(RepositoryRoot)src\TrainingArchive;$(RepositoryRoot)src\TransformerRanker;$(RepositoryRoot)external\eigen-3.4.0;%(AdditionalIncludeDirectories)</AdditionalIncludeDirectories>
<PreprocessorDefinitions>NOMINMAX;_CRT_SECURE_NO_WARNINGS;%(PreprocessorDefinitions)</PreprocessorDefinitions>
</ClCompile>
<Link>
<SubSystem>Console</SubSystem>
<GenerateDebugInformation>true</GenerateDebugInformation>
</Link>
</ItemDefinitionGroup>
<ItemGroup>
<ClCompile Include="TransformerRankerSelfTest.cpp" />
<ClCompile Include="..\..\src\TransformerRanker\TransformerRanker.cpp" />
</ItemGroup>
<ItemGroup>
<ClInclude Include="..\..\src\TransformerRanker\TransformerRanker.h" />
<ClInclude Include="..\..\src\TrainingArchive\TrainingArchive.h" />
</ItemGroup>
<Import Project="$(VCTargetsPath)\Microsoft.Cpp.targets" />
</Project>