Skip to content
Merged
Show file tree
Hide file tree
Changes from 9 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions lib/safetensors.ex
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,8 @@ defmodule Safetensors do
{:f, 64} => "F64",
{:f, 32} => "F32",
{:f, 16} => "F16",
{:f, 8} => "F8_E5M2",
{:f8_e4m3fn, 8} => "F8_E4M3",
{:s, 64} => "I64",
{:s, 32} => "I32",
{:s, 16} => "I16",
Expand All @@ -41,7 +43,7 @@ defmodule Safetensors do
{:u, 8} => "U8"
}

@dtype_to_type for {k, v} <- @type_to_dtype, into: %{}, do: {v, k}
@dtype_to_type for({k, v} <- @type_to_dtype, into: %{}, do: {v, k})

@doc """
Writes a map of tensors to a file.
Expand Down Expand Up @@ -94,8 +96,8 @@ defmodule Safetensors do
end

defp tensor_byte_size(tensor) do
{_, elem_size} = Nx.type(tensor)
elem_byte_size = div(elem_size, 8)
{_, size} = Nx.type(tensor)
elem_byte_size = div(size, 8)
Nx.size(tensor) * elem_byte_size
end

Expand Down
3 changes: 2 additions & 1 deletion mix.exs
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,8 @@ defmodule Safetensors.MixProject do
defp deps do
[
{:jason, "~> 1.4"},
{:nx, "~> 0.5"},
# TODO: Switch to released version once Nx with fp8 support is published
{:nx, github: "elixir-nx/nx", sparse: "nx", branch: "main"},
{:ex_doc, "~> 0.37", only: :dev, runtime: false}
]
end
Expand Down
2 changes: 1 addition & 1 deletion mix.lock
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,6 @@
"makeup_elixir": {:hex, :makeup_elixir, "1.0.1", "e928a4f984e795e41e3abd27bfc09f51db16ab8ba1aebdba2b3a575437efafc2", [:mix], [{:makeup, "~> 1.0", [hex: :makeup, repo: "hexpm", optional: false]}, {:nimble_parsec, "~> 1.2.3 or ~> 1.3", [hex: :nimble_parsec, repo: "hexpm", optional: false]}], "hexpm", "7284900d412a3e5cfd97fdaed4f5ed389b8f2b4cb49efc0eb3bd10e2febf9507"},
"makeup_erlang": {:hex, :makeup_erlang, "1.0.2", "03e1804074b3aa64d5fad7aa64601ed0fb395337b982d9bcf04029d68d51b6a7", [:mix], [{:makeup, "~> 1.0", [hex: :makeup, repo: "hexpm", optional: false]}], "hexpm", "af33ff7ef368d5893e4a267933e7744e46ce3cf1f61e2dccf53a111ed3aa3727"},
"nimble_parsec": {:hex, :nimble_parsec, "1.4.2", "8efba0122db06df95bfaa78f791344a89352ba04baedd3849593bfce4d0dc1c6", [:mix], [], "hexpm", "4b21398942dda052b403bbe1da991ccd03a053668d147d53fb8c4e0efe09c973"},
"nx": {:hex, :nx, "0.9.2", "17563029c01bf749aad3c31234326d7665abd0acc33ee2acbe531a4759f29a8a", [:mix], [{:complex, "~> 0.5", [hex: :complex, repo: "hexpm", optional: false]}, {:telemetry, "~> 0.4.0 or ~> 1.0", [hex: :telemetry, repo: "hexpm", optional: false]}], "hexpm", "914d74741617d8103de8ab1f8c880353e555263e1c397b8a1109f79a3716557f"},
"nx": {:git, "https://github.com/elixir-nx/nx.git", "04fe0ecf30cc20494f034f29fa3c07a3db7dd8c3", [sparse: "nx", branch: "main"]},
"telemetry": {:hex, :telemetry, "1.3.0", "fedebbae410d715cf8e7062c96a1ef32ec22e764197f70cda73d82778d61e7a2", [:rebar3], [], "hexpm", "7015fc8919dbe63764f4b4b87a95b7c0996bd539e0d499be6ec9d7f3875b79e6"},
}
142 changes: 142 additions & 0 deletions test/safetensors_test.exs
Original file line number Diff line number Diff line change
Expand Up @@ -78,4 +78,146 @@ defmodule SafetensorsTest do

assert Safetensors.load!(serialized) == %{"test1" => Nx.tensor([[0, 0], [0, 0]], type: :s32)}
end

describe "fp8 support" do
@tag :tmp_dir
test "write and read fp8 E4M3FN tensors", %{tmp_dir: tmp_dir} do
path = Path.join(tmp_dir, "fp8_e4m3fn")

# Create E4M3FN tensor
original = Nx.tensor([[1.0, 2.0], [3.0, 4.0]], type: :f8_e4m3fn)
Safetensors.write!(path, %{"weight" => original})

# Read back
loaded = Safetensors.read!(path)

# Verify type is preserved
assert Nx.type(loaded["weight"]) == {:f8_e4m3fn, 8}
assert Nx.shape(loaded["weight"]) == {2, 2}

# Note: Value accuracy testing is done in Nx core tests
# SafeTensors tests focus on type preservation and serialization
end

@tag :tmp_dir
test "write and read fp8 E5M2 tensors", %{tmp_dir: tmp_dir} do
path = Path.join(tmp_dir, "fp8_e5m2")

# Create E5M2 tensor
original = Nx.tensor([[1.0, 2.0], [3.0, 4.0]], type: :f8)
Safetensors.write!(path, %{"weight" => original})

# Read back
loaded = Safetensors.read!(path)

# Verify type is preserved
assert Nx.type(loaded["weight"]) == {:f, 8}
assert Nx.shape(loaded["weight"]) == {2, 2}
end

@tag :tmp_dir
test "round-trip preserves fp8 types", %{tmp_dir: tmp_dir} do
path = Path.join(tmp_dir, "fp8_mixed")

# Create tensors with different types
tensors = %{
"e4m3_weight" => Nx.tensor([[1.0, 2.0], [3.0, 4.0]], type: :f8_e4m3fn),
"e5m2_weight" => Nx.tensor([[5.0, 6.0], [7.0, 8.0]], type: :f8),
"f16_weight" => Nx.tensor([[9.0, 10.0], [11.0, 12.0]], type: :f16)
}

Safetensors.write!(path, tensors)
loaded = Safetensors.read!(path)

# Verify all types are preserved
assert Nx.type(loaded["e4m3_weight"]) == {:f8_e4m3fn, 8}
assert Nx.type(loaded["e5m2_weight"]) == {:f, 8}
assert Nx.type(loaded["f16_weight"]) == {:f, 16}
end

test "dump and load fp8 tensors" do
tensors = %{
"e4m3" => Nx.tensor([1.0, 2.0, 3.0], type: :f8_e4m3fn),
"e5m2" => Nx.tensor([4.0, 5.0, 6.0], type: :f8)
}

# Dump to binary
binary = tensors |> Safetensors.dump() |> IO.iodata_to_binary()

# Load back
loaded = Safetensors.load!(binary)

# Verify types
assert Nx.type(loaded["e4m3"]) == {:f8_e4m3fn, 8}
assert Nx.type(loaded["e5m2"]) == {:f, 8}

# Verify shapes
assert Nx.shape(loaded["e4m3"]) == {3}
assert Nx.shape(loaded["e5m2"]) == {3}
end

@tag :tmp_dir
test "lazy load fp8 tensors", %{tmp_dir: tmp_dir} do
path = Path.join(tmp_dir, "fp8_lazy")

# Write fp8 tensor
original = Nx.tensor([[1.0, 2.0], [3.0, 4.0]], type: :f8_e4m3fn)
Safetensors.write!(path, %{"weight" => original})

# Read lazily
%{"weight" => file_tensor} = Safetensors.read!(path, lazy: true)

# Verify it's a FileTensor
assert %Safetensors.FileTensor{} = file_tensor
assert file_tensor.type == {:f8_e4m3fn, 8}
assert file_tensor.shape == {2, 2}

# Convert to tensor and verify type is preserved
tensor = Nx.to_tensor(file_tensor)
assert Nx.type(tensor) == {:f8_e4m3fn, 8}
end

@tag :tmp_dir
test "fp8 tensor byte size calculation", %{tmp_dir: tmp_dir} do
path = Path.join(tmp_dir, "fp8_size")

# Create a large fp8 tensor
tensor = Nx.iota({100, 100}, type: :f8_e4m3fn)
Safetensors.write!(path, %{"large" => tensor})

# Verify file size is correct (8 bytes header size + header + 10000 bytes data)
file_size = File.stat!(path).size
header_start = 8

# Read header to get exact size
<<header_size::unsigned-64-integer-little, _rest::binary>> = File.read!(path)

# Data should be exactly 10000 bytes (100 * 100 * 1 byte per fp8)
expected_data_size = 10000
actual_data_size = file_size - header_start - header_size

assert actual_data_size == expected_data_size
end

test "fp8 dtype strings in header" do
# Create fp8 tensors
tensors = %{
"e4m3" => Nx.tensor([1.0], type: :f8_e4m3fn),
"e5m2" => Nx.tensor([2.0], type: :f8)
}

# Dump to binary
binary = tensors |> Safetensors.dump() |> IO.iodata_to_binary()

# Extract and parse header
<<header_size::unsigned-64-integer-little, header_json::binary-size(header_size),
_data::binary>> = binary

header = Jason.decode!(header_json)

# Verify dtype strings
assert header["e4m3"]["dtype"] == "F8_E4M3"
assert header["e5m2"]["dtype"] == "F8_E5M2"
end
end
end