diff --git a/deepseek/convert_dsv4.cc b/deepseek/convert_dsv4.cc index cfeeac3e..955dde45 100644 --- a/deepseek/convert_dsv4.cc +++ b/deepseek/convert_dsv4.cc @@ -19,11 +19,12 @@ // convert_dsv4 --weights --output // [--tokenizer_json ] [--fp4_high_first] +// [--mxfp4] [--lossless] [--verify_only] // // Handles: // * FP8 (e4m3fn) weights with 128x128-block e8m0 scales -> SFP. // * FP4 (e2m1, packed two per byte along the last dim, stored as I8) expert -// weights with per-32 e8m0 scales -> SFP. +// weights with per-32 e8m0 scales -> MXFP4 (with --mxfp4) or SFP. // * BF16/F32 tensors -> BF16/F32; I64 (hash routing table) -> F32. // * RoPE dim permutation: the reference applies rotary embeddings to // interleaved (even, odd) pairs, gemma.cpp to (i, i + dim/2) halves. @@ -115,6 +116,9 @@ void CompressAndWrite(const char* name, float* data, size_t rows, size_t cols, case Type::kF32: insert(float()); break; + case Type::kMXFP4: + insert(MxFp4Stream()); + break; default: HWY_ABORT("Unsupported output type %s for %s", TypeName(type), name); } @@ -305,7 +309,8 @@ struct ConvertArgs { bool verify_only = false; // With verify_only: check only the MTP tensors (fast pre-flight). bool mtp_only = false; - bool dspark = false; + bool mxfp4 = false; + bool lossless = false; }; class Converter { @@ -421,15 +426,93 @@ class Converter { return out; } - // Output type by source dtype: FP8/FP4 -> SFP, BF16 -> BF16, F32/I64 -> F32. + // Output type by source dtype: FP8/FP4 -> SFP (or MXFP4), BF16 -> BF16, F32/I64 -> F32. Type OutType(const std::string& src_name) { const SourceTensor* t = checkpoint_.Find(src_name); HWY_ASSERT_M(t != nullptr, src_name.c_str()); + if (args_.mxfp4 && t->dtype == "I8") return Type::kMXFP4; + if ((args_.lossless || args_.mxfp4) && t->dtype == "F8_E4M3") { + return Type::kBF16; + } if (t->dtype == "F8_E4M3" || t->dtype == "I8") return Type::kSFP; if (t->dtype == "BF16") return Type::kBF16; return Type::kF32; } + void WriteMxFp4Direct(const char* name, const std::string& src_name, + size_t rows, size_t cols) { + HWY_ASSERT(cols % 32 == 0); + const size_t num_blocks = (rows * cols) / 32; + const size_t packed_bytes = num_blocks * sizeof(MxFp4Stream); + + const SourceTensor* t = checkpoint_.Find(src_name); + HWY_ASSERT_M(t != nullptr, src_name.c_str()); + HWY_ASSERT(t->shape.size() == 2); + HWY_ASSERT(t->shape[0] == rows && t->shape[1] * 2 == cols); + + const E8M0Scale scale = LoadE8M0Scale(src_name); + const size_t scale_cols = cols / 32; + HWY_ASSERT(scale.tensor->shape[0] == rows && + scale.tensor->shape[1] == scale_cols); + + const std::vector raw = checkpoint_.ReadRaw(*t); + + if (args_.verify_only) { + ++num_written_; + if (num_written_ % 2000 == 0) { + fprintf(stderr, " ... %zu tensors verified\n", num_written_); + } + return; + } + + const Extents2D extents(rows, cols); + MatPtrT mat(name, extents); + mat.AppendTo(serialized_mat_ptrs_); + MatOwner owner; + owner.AllocateFor(mat, ctx_.allocator, MatPadding::kPacked); + HWY_ASSERT(mat.PackedBytes() == packed_bytes); + + uint8_t* dst = reinterpret_cast(mat.Packed()); + const size_t row_src_bytes = cols / 2; + const size_t row_blocks = cols / 32; + + for (size_t r = 0; r < rows; ++r) { + const uint8_t* src_row = raw.data() + r * row_src_bytes; + const uint8_t* srow = scale.raw.data() + r * scale_cols; + uint8_t* dst_row = dst + r * row_blocks * sizeof(MxFp4Stream); + + for (size_t b = 0; b < row_blocks; ++b) { + uint8_t* block = dst_row + b * sizeof(MxFp4Stream); + block[0] = srow[b]; + uint8_t* qs = block + 1; + const uint8_t* src_blk = src_row + b * 16; + uint8_t w[32]; + if (args_.fp4_high_first) { + for (size_t c2 = 0; c2 < 16; ++c2) { + const uint8_t byte = src_blk[c2]; + w[2 * c2] = byte >> 4; + w[2 * c2 + 1] = byte & 0x0F; + } + } else { + for (size_t c2 = 0; c2 < 16; ++c2) { + const uint8_t byte = src_blk[c2]; + w[2 * c2] = byte & 0x0F; + w[2 * c2 + 1] = byte >> 4; + } + } + for (size_t k = 0; k < 16; ++k) { + qs[k] = static_cast(w[k] | (w[k + 16] << 4)); + } + } + } + + writer_->Add(name, mat.Packed(), mat.PackedBytes()); + ++num_written_; + if (num_written_ % 500 == 0) { + fprintf(stderr, " ... %zu tensors written\n", num_written_); + } + } + void Write(const char* name, std::vector& data, size_t rows, size_t cols, Type type) { if (args_.verify_only) { @@ -489,9 +572,9 @@ class Converter { void Run() { const bool is_dspark = - args_.dspark || - (checkpoint_.Find("mtp.0.main_proj.weight") != nullptr); - ModelConfig config(Model::DEEPSEEK4_FLASH, Type::kSFP, + checkpoint_.Find("mtp.0.main_proj.weight") != nullptr; + const Type weight_type = args_.mxfp4 ? Type::kMXFP4 : Type::kSFP; + ModelConfig config(Model::DEEPSEEK4_FLASH, weight_type, PromptWrapping::GEMMA_IT); if (is_dspark) { config.num_mtp_layers = 3; @@ -731,12 +814,17 @@ class Converter { base.c_str()); } - if (checkpoint_.Find(src) == nullptr) { + const SourceTensor* src_t = checkpoint_.Find(src); + if (src_t == nullptr) { if (args_.verify_only) { ++num_skipped_; // shard not downloaded yet } return; } + if (args_.mxfp4 && src_t->dtype == "I8" && transform == 0) { + WriteMxFp4Direct(name.c_str(), src, rows, cols); + return; + } std::vector data = LoadF32(src, num); switch (transform) { case 1: @@ -807,8 +895,10 @@ int Main(int argc, char** argv) { args.verify_only = true; } else if (a == "--mtp_only") { args.mtp_only = true; - } else if (a == "--dspark") { - args.dspark = true; + } else if (a == "--mxfp4") { + args.mxfp4 = true; + } else if (a == "--lossless") { + args.lossless = true; } else { fprintf(stderr, "Unknown arg %s\n", a.c_str()); return 1; @@ -818,7 +908,7 @@ int Main(int argc, char** argv) { fprintf(stderr, "Usage: convert_dsv4 --weights --output " "[--tokenizer_json ] [--fp4_high_first] " - "[--verify_only]\n"); + "[--mxfp4] [--verify_only]\n"); return 1; } Converter(args).Run();