diff --git a/examples/fit-params/README.md b/examples/fit-params/README.md index 0f9506a1e1..3ca761ee6a 100644 --- a/examples/fit-params/README.md +++ b/examples/fit-params/README.md @@ -4,6 +4,9 @@ device memory, using measured metadata-only dry runs: the real generation pipeline is executed with graph building and memory measurement only, so no weight data is read and no ggml weight or compute buffers are allocated. +Safetensors inputs may contain only their complete header, including auxiliary +models and LoRAs. A successful fit does not verify that weight data is present; +generation still requires complete model files. Shaped host tensors are still materialized to carry state between graph builds; allocation failures are reported as fit errors. The measured per-module memory includes projected persistent cache buffers and is packed against the free diff --git a/src/core/util.cpp b/src/core/util.cpp index 4106cff52d..fb3e2c9d91 100644 --- a/src/core/util.cpp +++ b/src/core/util.cpp @@ -342,6 +342,7 @@ int32_t sd_get_num_physical_cores() { static sd_progress_cb_t sd_progress_cb = nullptr; void* sd_progress_cb_data = nullptr; static thread_local bool sd_progress_suppressed = false; +static thread_local bool sd_metadata_only_read = false; static sd_abort_cb_t sd_abort_cb = nullptr; static void* sd_abort_cb_data = nullptr; @@ -702,6 +703,19 @@ void sd_set_progress_suppressed(bool suppressed) { sd_progress_suppressed = suppressed; } +bool sd_get_metadata_only_read() { + return sd_metadata_only_read; +} + +SDMetadataOnlyReadScope::SDMetadataOnlyReadScope() + : previous_(sd_metadata_only_read) { + sd_metadata_only_read = true; +} + +SDMetadataOnlyReadScope::~SDMetadataOnlyReadScope() { + sd_metadata_only_read = previous_; +} + sd_image_t tensor_to_sd_image(const sd::Tensor& tensor, int frame_index) { const auto& shape = tensor.shape(); GGML_ASSERT(shape.size() == 4 || shape.size() == 5); diff --git a/src/core/util.h b/src/core/util.h index e0308d3309..dadac80f83 100644 --- a/src/core/util.h +++ b/src/core/util.h @@ -94,6 +94,18 @@ void* sd_get_progress_callback_data(); bool sd_get_progress_suppressed(); void sd_set_progress_suppressed(bool suppressed); +bool sd_get_metadata_only_read(); + +class SDMetadataOnlyReadScope { + bool previous_; + +public: + SDMetadataOnlyReadScope(); + ~SDMetadataOnlyReadScope(); + SDMetadataOnlyReadScope(const SDMetadataOnlyReadScope&) = delete; + SDMetadataOnlyReadScope& operator=(const SDMetadataOnlyReadScope&) = delete; +}; + sd_preview_cb_t sd_get_preview_callback(); void* sd_get_preview_callback_data(); preview_t sd_get_preview_mode(); diff --git a/src/model/adapter/lora.hpp b/src/model/adapter/lora.hpp index 6c0f41a0aa..9c19e46996 100644 --- a/src/model/adapter/lora.hpp +++ b/src/model/adapter/lora.hpp @@ -131,7 +131,7 @@ struct LoraModel : public GGMLRunner { for (const auto& pair : lora_tensors) { lora_params.push_back(pair.second); } - if (!model_manager->prepare_params(lora_params)) { + if (!GGMLRunner::measure_mode_enabled() && !model_manager->prepare_params(lora_params)) { LOG_ERROR("lora model manager prepare params failed"); return false; } @@ -409,6 +409,11 @@ struct LoraModel : public GGMLRunner { return compatible; } + static float read_scale(ggml_tensor* tensor) { + // Scalar values do not affect graph shapes during memory measurement. + return GGMLRunner::measure_mode_enabled() ? 1.0f : ggml_ext_backend_tensor_get_f32(tensor); + } + ggml_tensor* get_lora_weight_diff(const std::string& model_tensor_name, ggml_context* ctx, ggml_backend_t backend) { ggml_tensor* updown = nullptr; int index = 0; @@ -461,12 +466,12 @@ struct LoraModel : public GGMLRunner { int64_t rank = lora_down->ne[ggml_n_dims(lora_down) - 1]; iter = lora_tensors.find(scale_name); if (iter != lora_tensors.end()) { - scale_value = ggml_ext_backend_tensor_get_f32(iter->second); + scale_value = read_scale(iter->second); applied_lora_tensors.insert(scale_name); } else { iter = lora_tensors.find(alpha_name); if (iter != lora_tensors.end()) { - float alpha = ggml_ext_backend_tensor_get_f32(iter->second); + float alpha = read_scale(iter->second); scale_value = alpha / rank; // LOG_DEBUG("rank %s %ld %.2f %.2f", alpha_name.c_str(), rank, alpha, scale_value); applied_lora_tensors.insert(alpha_name); @@ -615,7 +620,7 @@ struct LoraModel : public GGMLRunner { int64_t rank = hada_1_down->ne[ggml_n_dims(hada_1_down) - 1]; iter = lora_tensors.find(alpha_name); if (iter != lora_tensors.end()) { - float alpha = ggml_ext_backend_tensor_get_f32(iter->second); + float alpha = read_scale(iter->second); scale_value = alpha / rank; applied_lora_tensors.insert(alpha_name); } @@ -728,7 +733,7 @@ struct LoraModel : public GGMLRunner { float scale_value = 1.0f; iter = lora_tensors.find(alpha_name); if (iter != lora_tensors.end()) { - float alpha = ggml_ext_backend_tensor_get_f32(iter->second); + float alpha = read_scale(iter->second); scale_value = alpha / rank; applied_lora_tensors.insert(alpha_name); } @@ -889,7 +894,7 @@ struct LoraModel : public GGMLRunner { float scale_value = 1.0f; iter = lora_tensors.find(alpha_name); if (iter != lora_tensors.end()) { - float alpha = ggml_ext_backend_tensor_get_f32(iter->second); + float alpha = read_scale(iter->second); scale_value = alpha / rank; } @@ -1016,12 +1021,12 @@ struct LoraModel : public GGMLRunner { int64_t rank = lora_down->ne[ggml_n_dims(lora_down) - 1]; iter = lora_tensors.find(scale_name); if (iter != lora_tensors.end()) { - scale_value = ggml_ext_backend_tensor_get_f32(iter->second); + scale_value = read_scale(iter->second); scale_tensor_name = scale_name; } else { iter = lora_tensors.find(alpha_name); if (iter != lora_tensors.end()) { - float alpha = ggml_ext_backend_tensor_get_f32(iter->second); + float alpha = read_scale(iter->second); scale_value = alpha / rank; scale_tensor_name = alpha_name; // LOG_DEBUG("rank %s %ld %.2f %.2f", alpha_name.c_str(), rank, alpha, scale_value); diff --git a/src/model_io/safetensors_io.cpp b/src/model_io/safetensors_io.cpp index df71eab11a..f63e4ded9b 100644 --- a/src/model_io/safetensors_io.cpp +++ b/src/model_io/safetensors_io.cpp @@ -5,6 +5,7 @@ #include #include #include +#include #include #include #include @@ -180,9 +181,19 @@ bool read_safetensors_file(const std::string& file_path, continue; } - size_t begin = tensor_info["data_offsets"][0].get(); - size_t end = tensor_info["data_offsets"][1].get(); - if (begin > end || end > file_size_ - data_start) { + const auto& offsets = tensor_info["data_offsets"]; + const auto valid_offset = [](const nlohmann::json& value) { + return value.is_number_unsigned() ? value.get() <= std::numeric_limits::max() + : value.is_number_integer() && value.get() >= 0; + }; + if (!offsets.is_array() || offsets.size() != 2 || !valid_offset(offsets[0]) || !valid_offset(offsets[1])) { + set_error(error, "invalid data offsets for tensor '" + name + "'"); + return false; + } + size_t begin = offsets[0].get(); + size_t end = offsets[1].get(); + if (begin > end || end > std::numeric_limits::max() - data_start || + (!sd_get_metadata_only_read() && end > file_size_ - data_start)) { set_error(error, "data offsets out of bounds for tensor '" + name + "'"); return false; } @@ -193,15 +204,30 @@ bool read_safetensors_file(const std::string& file_path, return false; } - if (shape.size() > SD_MAX_DIMS) { + if (!shape.is_array() || shape.size() > SD_MAX_DIMS) { set_error(error, "invalid tensor '" + name + "'"); return false; } int n_dims = (int)shape.size(); int64_t ne[SD_MAX_DIMS] = {1, 1, 1, 1, 1}; + // Bound intermediate products too, including F64/I64 conversion and zero-sized tensors. + const uint64_t max_nelements = + std::min(std::numeric_limits::max(), std::numeric_limits::max()) / + (2 * ggml_type_size(type)); + uint64_t nelements = 1; for (int i = 0; i < n_dims; i++) { - ne[i] = shape[i].get(); + if (!shape[i].is_number_unsigned()) { + set_error(error, "invalid shape for tensor '" + name + "'"); + return false; + } + const uint64_t dim = shape[i].get(); + if (dim > max_nelements || (dim != 0 && nelements > max_nelements / dim)) { + set_error(error, "invalid shape for tensor '" + name + "'"); + return false; + } + nelements *= std::max(dim, 1); + ne[i] = static_cast(dim); } if (n_dims == 5) { @@ -226,11 +252,11 @@ bool read_safetensors_file(const std::string& file_path, if (dtype == "F8_E4M3") { tensor_storage.is_f8_e4m3 = true; // f8 -> f16 - tensor_size_ok = (tensor_storage.nbytes() == tensor_data_size * 2); + tensor_size_ok = (tensor_storage.nbytes() / 2 == tensor_data_size); } else if (dtype == "F8_E5M2") { tensor_storage.is_f8_e5m2 = true; // f8 -> f16 - tensor_size_ok = (tensor_storage.nbytes() == tensor_data_size * 2); + tensor_size_ok = (tensor_storage.nbytes() / 2 == tensor_data_size); } else if (dtype == "F64") { tensor_storage.is_f64 = true; // f64 -> f32 diff --git a/src/stable-diffusion.cpp b/src/stable-diffusion.cpp index f0363d5d56..2ec7069573 100644 --- a/src/stable-diffusion.cpp +++ b/src/stable-diffusion.cpp @@ -4054,6 +4054,7 @@ enum sd_fit_status_t sd_fit_params(const sd_ctx_params_t* sd_ctx_params, return SD_FIT_FAILURE; } + ggml_time_init(); int64_t t0 = ggml_time_ms(); sd_ctx_params_t dry_params = *sd_ctx_params; @@ -4064,7 +4065,14 @@ enum sd_fit_status_t sd_fit_params(const sd_ctx_params_t* sd_ctx_params, sd_ctx_t* sd_ctx = &sd_ctx_storage; sd_ctx->sd = new StableDiffusionGGML(); sd_ctx->sd->fit_dry_run = true; - if (!sd_ctx->sd->init(&dry_params)) { + bool init_ok = false; + try { + SDMetadataOnlyReadScope metadata_only; + init_ok = sd_ctx->sd->init(&dry_params); + } catch (const std::exception& error) { + LOG_ERROR("fit-params: dry-run model init failed: %s", error.what()); + } + if (!init_ok) { LOG_ERROR("fit-params: dry-run model init failed"); delete sd_ctx->sd; sd_ctx->sd = nullptr; @@ -4128,6 +4136,7 @@ enum sd_fit_status_t sd_fit_params(const sd_ctx_params_t* sd_ctx_params, auto measure = [&](const sd_tiling_params_t& tiling, std::vector& records) -> bool { + SDMetadataOnlyReadScope metadata_only; records.clear(); struct MeasureModeGuard { explicit MeasureModeGuard(std::vector* records) { diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 6c1ebe2c95..11050ef9e9 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -11,6 +11,11 @@ target_include_directories(test-fit-params PRIVATE "${PROJECT_SOURCE_DIR}/src") target_link_libraries(test-fit-params PRIVATE stable-diffusion ${CMAKE_THREAD_LIBS_INIT}) add_test(NAME test-fit-params COMMAND test-fit-params) +add_executable(test-safetensors-metadata test_safetensors_metadata.cpp) +target_include_directories(test-safetensors-metadata PRIVATE "${PROJECT_SOURCE_DIR}/src") +target_link_libraries(test-safetensors-metadata PRIVATE stable-diffusion ${CMAKE_THREAD_LIBS_INIT}) +add_test(NAME test-safetensors-metadata COMMAND test-safetensors-metadata) + add_executable(test-ltx-vae-temporal test-ltx-vae-temporal.cpp) target_link_libraries(test-ltx-vae-temporal PRIVATE stable-diffusion) add_test(NAME test-ltx-vae-temporal COMMAND test-ltx-vae-temporal) diff --git a/tests/test_safetensors_metadata.cpp b/tests/test_safetensors_metadata.cpp new file mode 100644 index 0000000000..3782ab4714 --- /dev/null +++ b/tests/test_safetensors_metadata.cpp @@ -0,0 +1,268 @@ +#include "core/util.h" +#include "model/adapter/lora.hpp" +#include "model_io/binary_io.h" +#include "model_io/safetensors_io.h" + +#include +#include +#include +#include +#include + +namespace safetensors_metadata_test { + + const std::string header = + R"({"model.diffusion_model.test.weight.diff":{"dtype":"F32","shape":[6,4],"data_offsets":[0,96]}})"; + + void write_file(const std::filesystem::path& path, const std::string& json, size_t payload = 0) { + std::ofstream file(path, std::ios::binary | std::ios::trunc); + GGML_ASSERT(file.is_open()); + model_io::write_u64(file, json.size()); + file.write(json.data(), static_cast(json.size())); + for (size_t i = 0; i < payload; ++i) { + file.put('\0'); + } + GGML_ASSERT(file.good()); + } + + bool read_file(const std::filesystem::path& path) { + std::vector tensors; + std::string error; + return read_safetensors_file(path.string(), tensors, &error); + } + + void test_reader(const std::filesystem::path& path) { + write_file(path, header, 96); + std::vector full; + GGML_ASSERT(read_safetensors_file(path.string(), full)); + GGML_ASSERT(full.size() == 1); + write_file(path, header); + GGML_ASSERT(is_safetensors_file(path.string())); + GGML_ASSERT(!read_file(path)); + { + SDMetadataOnlyReadScope scope; + std::vector metadata; + GGML_ASSERT(read_safetensors_file(path.string(), metadata)); + GGML_ASSERT(metadata.size() == full.size()); + GGML_ASSERT(metadata[0].to_string() == full[0].to_string()); + GGML_ASSERT(metadata[0].nbytes() == 96); + { + SDMetadataOnlyReadScope nested; + GGML_ASSERT(read_file(path)); + } + GGML_ASSERT(read_file(path)); + bool other_thread_accepted = true; + std::thread other([&]() { other_thread_accepted = read_file(path); }); + other.join(); + GGML_ASSERT(!other_thread_accepted); + } + GGML_ASSERT(!read_file(path)); + write_file(path, header, 4); + GGML_ASSERT(!read_file(path)); + { + SDMetadataOnlyReadScope scope; + GGML_ASSERT(read_file(path)); + } + } + + void test_invalid_headers(const std::filesystem::path& path) { + SDMetadataOnlyReadScope scope; + for (const auto& json : { + R"({"x":{"dtype":"F32","shape":[-1000000],"data_offsets":[0,18446744073705551616]}})", + R"({"x":{"dtype":"F32","shape":[2305843009213693952],"data_offsets":[0,9223372036854775808]}})", + R"({"x":{"dtype":"F32","shape":[4294967296,4294967296,1,1,1],"data_offsets":[0,0]}})", + R"({"x":{"dtype":"F32","shape":[0,4294967296,4294967296,1,1],"data_offsets":[0,0]}})", + R"({"x":{"dtype":"F64","shape":[2305843009213693952],"data_offsets":[0,0]}})", + R"({"x":{"dtype":"I64","shape":[2305843009213693952],"data_offsets":[0,0]}})", + R"({"x":{"dtype":"F32","shape":[18446744073709551615],"data_offsets":[0,0]}})", + R"({"x":{"dtype":"F32","shape":[4.0],"data_offsets":[0,16]}})", + R"({"x":{"dtype":"F32","shape":[true],"data_offsets":[0,4]}})", + R"({"x":{"dtype":"F32","shape":null,"data_offsets":[0,4]}})", + R"({"x":{"dtype":"F32","shape":[1],"data_offsets":[4,0]}})", + R"({"x":{"dtype":"F32","shape":[1],"data_offsets":[-1000,-996]}})", + R"({"x":{"dtype":"F32","shape":[1],"data_offsets":[0.0,4.0]}})", + R"({"x":{"dtype":"F32","shape":[1],"data_offsets":[0]}})", + R"({"x":{"dtype":"F32","shape":[1],"data_offsets":[0,8]}})", + R"({"x":{"dtype":"U16","shape":[1],"data_offsets":[0,2]}})", + R"({"x":{"dtype":"F32","shape":[1,1,1,1,1,1],"data_offsets":[0,4]}})", + "{invalid json"}) { + write_file(path, json); + GGML_ASSERT(!read_file(path)); + } + const auto max_offset = std::numeric_limits::max(); + write_file(path, "{\"x\":{\"dtype\":\"F32\",\"shape\":[1],\"data_offsets\":[" + + std::to_string(max_offset - 4) + "," + std::to_string(max_offset) + "]}}"); + GGML_ASSERT(!read_file(path)); + write_file(path, "{\"x\":{\"dtype\":\"F8_E4M3\",\"shape\":[1],\"data_offsets\":[0," + + std::to_string(max_offset / 2 + 2) + "]}}"); + GGML_ASSERT(!read_file(path)); + write_file(path, header); + std::filesystem::resize_file(path, 8 + header.size() - 1); + GGML_ASSERT(!read_file(path)); + } + + void test_converted_types(const std::filesystem::path& path) { + SDMetadataOnlyReadScope scope; + for (const auto& json : { + R"({"x":{"dtype":"F8_E4M3","shape":[2],"data_offsets":[0,2]}})", + R"({"x":{"dtype":"F8_E5M2","shape":[2],"data_offsets":[0,2]}})", + R"({"x":{"dtype":"F64","shape":[2],"data_offsets":[0,16]}})", + R"({"x":{"dtype":"I64","shape":[2],"data_offsets":[0,16]}})", + R"({"x":{"dtype":"F16","shape":[2],"data_offsets":[0,4]}})"}) { + write_file(path, json); + GGML_ASSERT(read_file(path)); + } + } + + void test_shapes(const std::filesystem::path& path) { + for (const std::string shape : {"[]", "[2,3,4,1,1]", "[0,2,3]", "[2,0,3]", "[2,3,0]"}) { + const size_t payload = shape == "[]" ? 4 : shape == "[2,3,4,1,1]" ? 96 + : 0; + const std::string json = "{\"x\":{\"dtype\":\"F32\",\"shape\":" + shape + + ",\"data_offsets\":[0," + std::to_string(payload) + "]}}"; + write_file(path, json, payload); + GGML_ASSERT(read_file(path)); + write_file(path, json); + SDMetadataOnlyReadScope scope; + GGML_ASSERT(read_file(path)); + } + for (const std::string shape : {"[-1,0]", "[4.0]", "[0,4294967296,4294967296,1,1]"}) { + write_file(path, "{\"x\":{\"dtype\":\"F32\",\"shape\":" + shape + ",\"data_offsets\":[0,0]}}"); + GGML_ASSERT(!read_file(path)); + sd_ctx_params_t params; + sd_ctx_params_init(¶ms); + const auto model_path = path.string(); + params.model_path = model_path.c_str(); + sd_fit_workload_t workload; + sd_fit_workload_init(&workload); + sd_fit_result_t result{}; + GGML_ASSERT(sd_fit_params(¶ms, &workload, &result) == SD_FIT_ERROR); + sd_fit_result_free(&result); + GGML_ASSERT(!sd_get_metadata_only_read()); + } + } + + void test_shard_and_lora(const std::filesystem::path& path) { + write_file(path, header); + auto index = path; + index += ".index.json"; + { + std::ofstream file(index); + file << "{\"weight_map\":{\"model.diffusion_model.test.weight.diff\":\"" + << path.filename().string() << "\"}}"; + } + ModelLoader strict; + GGML_ASSERT(!strict.init_from_file(index.string())); + ggml_backend_t cpu = sd_backend_cpu_init(); + GGML_ASSERT(cpu != nullptr); + { + SDMetadataOnlyReadScope scope; + ModelLoader loader; + GGML_ASSERT(loader.init_from_file(index.string())); + GGML_ASSERT(loader.get_tensor_storage_map().size() == 1); + GGMLRunner::set_measure_mode(true); + { + LoraModel lora("metadata", cpu, cpu, path.string()); + GGML_ASSERT(lora.load_from_file(1)); + GGML_ASSERT(!lora.lora_tensors.empty()); + } + GGMLRunner::set_measure_mode(false); + } + { + LoraModel lora("strict", cpu, cpu, path.string()); + GGML_ASSERT(!lora.load_from_file(1)); + } + ggml_backend_free(cpu); + std::filesystem::remove(index); + } + + void test_lora_scalars(const std::filesystem::path& path) { + ggml_backend_t cpu = sd_backend_cpu_init(); + GGML_ASSERT(cpu != nullptr); + for (const std::string suffix : {"alpha", "scale"}) { + const std::string json = + R"({"model.diffusion_model.test.lora_down.weight":{"dtype":"F32","shape":[2,4],"data_offsets":[0,32]},)" + R"("model.diffusion_model.test.lora_up.weight":{"dtype":"F32","shape":[6,2],"data_offsets":[32,80]},)" + "\"model.diffusion_model.test." + + suffix + R"(":{"dtype":"F32","shape":[1],"data_offsets":[80,84]}})"; + int full_nodes = 0; + for (bool measure : {false, true}) { + write_file(path, json, measure ? 0 : 84); + SDMetadataOnlyReadScope scope; + GGMLRunner::set_measure_mode(measure); + { + LoraModel lora("scalars", cpu, cpu, path.string()); + GGML_ASSERT(lora.load_from_file(1)); + auto* scalar = lora.lora_tensors.at("lora.model.diffusion_model.test.weight." + suffix); + GGML_ASSERT(LoraModel::read_scale(scalar) == (measure ? 1.0f : 0.0f)); + GGML_ASSERT((scalar->data == nullptr) == measure); + ggml_init_params init = {}; + init.mem_size = 128 * ggml_tensor_overhead() + ggml_graph_overhead_custom(128, false); + init.no_alloc = true; + auto* ctx = ggml_init(init); + GGML_ASSERT(ctx != nullptr); + ggml_set_name(ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 1), "ggml_runner_build_in_tensor:one"); + auto* weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 4, 6); + auto* diff = lora.get_lora_weight_diff("model.diffusion_model.test.weight", ctx, cpu); + GGML_ASSERT(diff != nullptr && ggml_are_same_shape(diff, weight)); + auto* x = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 4, 1); + WeightAdapter::ForwardParams params = {}; + params.op_type = WeightAdapter::ForwardParams::op_type_t::OP_LINEAR; + auto* out = lora.get_out_diff(ctx, cpu, x, weight, params, "model.diffusion_model.test.weight"); + GGML_ASSERT(out != nullptr && out->ne[0] == 6 && out->ne[1] == 1); + auto* graph = ggml_new_graph_custom(ctx, 128, false); + ggml_build_forward_expand(graph, diff); + ggml_build_forward_expand(graph, out); + if (measure) { + GGML_ASSERT(ggml_graph_n_nodes(graph) == full_nodes); + } else { + full_nodes = ggml_graph_n_nodes(graph); + } + ggml_free(ctx); + } + GGMLRunner::set_measure_mode(false); + } + } + ggml_backend_free(cpu); + } + + void test_exception_cleanup(const std::filesystem::path& path) { + write_file(path, R"({"x":{"dtype":42,"shape":[1],"data_offsets":[0,4]}})"); + bool caught = false; + try { + SDMetadataOnlyReadScope scope; + read_file(path); + } catch (const std::exception&) { + caught = true; + } + GGML_ASSERT(caught); + GGML_ASSERT(!sd_get_metadata_only_read()); + sd_ctx_params_t params; + sd_ctx_params_init(¶ms); + const auto model_path = path.string(); + params.model_path = model_path.c_str(); + sd_fit_workload_t workload; + sd_fit_workload_init(&workload); + sd_fit_result_t result{}; + GGML_ASSERT(sd_fit_params(¶ms, &workload, &result) == SD_FIT_ERROR); + sd_fit_result_free(&result); + GGML_ASSERT(!sd_get_metadata_only_read()); + write_file(path, header); + GGML_ASSERT(!read_file(path)); + } + +} // namespace safetensors_metadata_test + +int main() { + using namespace safetensors_metadata_test; + const auto path = std::filesystem::temp_directory_path() / "sd-test-safetensors-metadata.safetensors"; + test_exception_cleanup(path); + test_reader(path); + test_invalid_headers(path); + test_converted_types(path); + test_shapes(path); + test_shard_and_lora(path); + test_lora_scalars(path); + std::filesystem::remove(path); + return 0; +}