From da3abd0089721c0286998d71e6784d8190029d66 Mon Sep 17 00:00:00 2001 From: Daniel Kukula Date: Wed, 30 Sep 2026 22:07:56 +0200 Subject: [PATCH] Add Bumblebee.load_embedding_head for SentenceTransformers pipelines --- lib/bumblebee.ex | 80 ++++- lib/bumblebee/huggingface/hub.ex | 2 +- .../huggingface/sentence_transformers.ex | 307 ++++++++++++++++++ lib/bumblebee/text/text_embedding.ex | 3 + .../sentence_transformers_test.exs | 208 ++++++++++++ 5 files changed, 586 insertions(+), 14 deletions(-) create mode 100644 lib/bumblebee/huggingface/sentence_transformers.ex create mode 100644 test/bumblebee/huggingface/sentence_transformers_test.exs diff --git a/lib/bumblebee.ex b/lib/bumblebee.ex index 28314c44..11af3b28 100644 --- a/lib/bumblebee.ex +++ b/lib/bumblebee.ex @@ -640,6 +640,62 @@ defmodule Bumblebee do end end + @doc """ + Loads an embedding head from a SentenceTransformers model repository and attaches + it to the given model. + + Reads `modules.json` from the repository and stacks the corresponding layers + (such as pooling, dense projections, and normalization) on top of `model_info.model`, + and loads the corresponding parameters into `model_info.params`. + + The final output of the resulting model is a map with the `:embedding` key. + + ## Options + + * `:safetensors_reader` - a function that reads a safetensors file into + parameters map. Defaults to `Safetensors.read!/2` with `lazy: true` + + * `:backend` - the backend to allocate the tensors on. It is either + an atom or a tuple in the shape `{backend, options}` + + * `:type` - either a type or `Axon.MixedPrecision` policy to apply + to the model parameters + + ## Examples + + {:ok, model_info} = Bumblebee.load_model({:hf, "google/embeddinggemma-300m"}) + {:ok, model_info} = Bumblebee.load_embedding_head({:hf, "google/embeddinggemma-300m"}, model_info) + serving = Bumblebee.Text.text_embedding(model_info, tokenizer) + + """ + @doc type: :model + @spec load_embedding_head(repository(), model_info(), keyword()) :: + {:ok, model_info()} | {:error, String.t()} + def load_embedding_head(repository, model_info, opts \\ []) do + repository = normalize_repository!(repository) + + opts = + Keyword.validate!(opts, [ + :safetensors_reader, + :backend, + :type + ]) + + case get_repo_files(repository) do + {:ok, repo_files} -> + HuggingFace.SentenceTransformers.load_embedding_head( + repository, + repo_files, + &download(repository, &1, repo_files[&1]), + model_info, + opts + ) + + {:error, message} -> + {:error, message} + end + end + defp maybe_load_model_spec(opts, repository, repo_files) do spec_result = if spec = opts[:spec] do @@ -1274,19 +1330,17 @@ defmodule Bumblebee do end defp get_repo_files({:local, dir}) do - case File.ls(dir) do - {:ok, filenames} -> - repo_files = - for filename <- filenames, - path = Path.join(dir, filename), - File.regular?(path), - into: %{}, - do: {filename, nil} - - {:ok, repo_files} - - {:error, reason} -> - {:error, "could not read #{dir}, reason: #{:file.format_error(reason)}"} + if File.dir?(dir) do + repo_files = + dir + |> Path.join("**") + |> Path.wildcard(match_dot: true) + |> Enum.filter(&File.regular?/1) + |> Map.new(fn path -> {Path.relative_to(path, dir), nil} end) + + {:ok, repo_files} + else + {:error, "could not read #{dir}, reason: no such file or directory"} end end diff --git a/lib/bumblebee/huggingface/hub.ex b/lib/bumblebee/huggingface/hub.ex index 6a730301..312a2af6 100644 --- a/lib/bumblebee/huggingface/hub.ex +++ b/lib/bumblebee/huggingface/hub.ex @@ -21,7 +21,7 @@ defmodule Bumblebee.HuggingFace.Hub do def file_listing_url(repository_id, subdir, revision) do revision = revision || "main" path = if(subdir, do: "/" <> subdir) - @huggingface_endpoint <> "/api/models/#{repository_id}/tree/#{revision}#{path}" + @huggingface_endpoint <> "/api/models/#{repository_id}/tree/#{revision}#{path}?recursive=true" end @doc """ diff --git a/lib/bumblebee/huggingface/sentence_transformers.ex b/lib/bumblebee/huggingface/sentence_transformers.ex new file mode 100644 index 00000000..0d27167a --- /dev/null +++ b/lib/bumblebee/huggingface/sentence_transformers.ex @@ -0,0 +1,307 @@ +defmodule Bumblebee.HuggingFace.SentenceTransformers do + @moduledoc false + + alias Bumblebee.Layers + + @doc """ + Loads a SentenceTransformers embedding head and attaches it to the model. + """ + def load_embedding_head(_repository, repo_files, download_fun, model_info, opts) do + case repo_files do + %{"modules.json" => _etag} -> + with {:ok, modules_path} <- download_fun.("modules.json"), + {:ok, modules} <- decode_json(modules_path) do + modules = Enum.sort_by(modules, & &1["idx"]) + + case modules do + [%{"type" => "sentence_transformers.models.Transformer"} | remaining] -> + build_pipeline(remaining, repo_files, download_fun, model_info, opts) + + _other -> + {:error, + "expected the first module in modules.json to be sentence_transformers.models.Transformer"} + end + end + + _ -> + {:error, "could not find modules.json in the repository"} + end + end + + defp build_pipeline(modules, repo_files, download_fun, model_info, opts) do + base_model = model_info.model + hidden_state = Axon.nx(base_model, & &1.hidden_state) + + attention_mask = + Layers.default Axon.input("attention_mask", optional: true) do + Layers.default_attention_mask(Axon.input("input_ids")) + end + + initial_state = {hidden_state, model_info.params.data} + + result = + Enum.reduce_while(modules, {:ok, initial_state}, fn module, {:ok, {node, params_data}} -> + case build_module( + module, + node, + attention_mask, + params_data, + repo_files, + download_fun, + opts + ) do + {:ok, node, params_data} -> + {:cont, {:ok, {node, params_data}}} + + {:error, reason} -> + {:halt, {:error, reason}} + end + end) + + with {:ok, {embedding, params_data}} <- result do + final_model = Axon.container(%{embedding: embedding, hidden_state: hidden_state}) + final_model = apply_type(final_model, opts[:type]) + final_params = %{model_info.params | data: params_data} + + {:ok, %{model_info | model: final_model, params: final_params}} + end + end + + defp build_module( + %{"type" => "sentence_transformers.models.Pooling", "path" => path}, + node, + attention_mask, + params_data, + _repo_files, + download_fun, + _opts + ) do + config_file = Path.join(path, "config.json") + + with {:ok, config_path} <- download_fun.(config_file), + {:ok, config} <- decode_json(config_path) do + node = pooling_layer(node, attention_mask, config, path) + {:ok, node, params_data} + end + end + + defp build_module( + %{"type" => "sentence_transformers.models.Dense", "path" => path}, + node, + _attention_mask, + params_data, + repo_files, + download_fun, + opts + ) do + config_file = Path.join(path, "config.json") + + with {:ok, config_path} <- download_fun.(config_file), + {:ok, config} <- decode_json(config_path), + {:ok, weights_file} <- find_weights_file(repo_files, path), + {:ok, weights_path} <- download_fun.(weights_file) do + out_features = config["out_features"] + bias = Map.get(config, "bias", true) + activation = parse_activation(config["activation_function"]) + + layer_name = path + + node = + node + |> Axon.dense(out_features, use_bias: bias, name: layer_name) + |> maybe_activation(activation) + + tensors = load_tensors(weights_file, weights_path, opts) + + kernel = + tensors["linear.weight"] + |> Nx.to_tensor() + |> Nx.transpose() + |> cast_param(opts[:type]) + |> allocate_param(opts[:backend]) + + layer_params = + if bias do + bias_tensor = + tensors["linear.bias"] + |> Nx.to_tensor() + |> cast_param(opts[:type]) + |> allocate_param(opts[:backend]) + + %{"kernel" => kernel, "bias" => bias_tensor} + else + %{"kernel" => kernel} + end + + params_data = Map.put(params_data, layer_name, layer_params) + {:ok, node, params_data} + end + end + + defp build_module( + %{"type" => "sentence_transformers.models.Normalize", "path" => path}, + node, + _attention_mask, + params_data, + _repo_files, + _download_fun, + _opts + ) do + node = Axon.nx(node, &Bumblebee.Utils.Nx.normalize/1, name: "#{path}.normalize") + {:ok, node, params_data} + end + + defp build_module( + %{"type" => "sentence_transformers.models.Dropout"}, + node, + _attention_mask, + params_data, + _repo_files, + _download_fun, + _opts + ) do + {:ok, node, params_data} + end + + defp build_module(%{"type" => type}, _node, _mask, _params, _repo_files, _download_fun, _opts) do + {:error, "unsupported SentenceTransformers module #{inspect(type)}"} + end + + defp pooling_layer(hidden_state, attention_mask, config, name) do + modes = + for {key, mode} <- [ + {"pooling_mode_cls_token", :cls_token}, + {"pooling_mode_mean_tokens", :mean_tokens}, + {"pooling_mode_max_tokens", :max_tokens}, + {"pooling_mode_mean_sqrt_len_tokens", :mean_sqrt_len_tokens}, + {"pooling_mode_lasttoken", :last_token} + ], + config[key] == true, + do: mode + + Axon.layer( + fn hidden_state, attention_mask, _opts -> + pool_outputs = + Enum.map(modes, fn + :mean_tokens -> + mask = Nx.new_axis(attention_mask, -1) + sum_embeddings = Nx.sum(Nx.multiply(hidden_state, mask), axes: [1]) + sum_mask = Nx.sum(mask, axes: [1]) |> Nx.max(1.0e-9) + Nx.divide(sum_embeddings, sum_mask) + + :cls_token -> + hidden_state[[.., 0, ..]] + + :max_tokens -> + mask = Nx.new_axis(attention_mask, -1) + + hidden_state + |> Nx.select(mask, Nx.Constants.min_finite(hidden_state)) + |> Nx.reduce_max(axes: [1]) + + :mean_sqrt_len_tokens -> + mask = Nx.new_axis(attention_mask, -1) + sum_embeddings = Nx.sum(Nx.multiply(hidden_state, mask), axes: [1]) + sum_mask = Nx.sum(mask, axes: [1]) |> Nx.max(1.0e-9) + Nx.divide(sum_embeddings, Nx.sqrt(sum_mask)) + + :last_token -> + lengths = + attention_mask + |> Nx.sum(axes: [1]) + |> Nx.subtract(1) + |> Nx.as_type({:s, 64}) + + Bumblebee.Utils.Nx.batched_take(hidden_state, lengths) + end) + + case pool_outputs do + [single] -> single + multiple -> Nx.concatenate(multiple, axis: -1) + end + end, + [hidden_state, attention_mask], + name: name + ) + end + + defp parse_activation("torch.nn.modules.linear.Identity"), do: nil + defp parse_activation("torch.nn.modules.activation.Tanh"), do: :tanh + defp parse_activation("torch.nn.modules.activation.ReLU"), do: :relu + defp parse_activation("torch.nn.modules.activation.GELU"), do: :gelu + defp parse_activation("torch.nn.modules.activation.SiLU"), do: :silu + defp parse_activation(nil), do: nil + defp parse_activation(other), do: raise("unsupported activation function #{inspect(other)}") + + defp maybe_activation(node, nil), do: node + defp maybe_activation(node, activation), do: Axon.activation(node, activation) + + defp find_weights_file(repo_files, dir) do + safetensors = Path.join(dir, "model.safetensors") + pytorch_bin = Path.join(dir, "pytorch_model.bin") + + cond do + Map.has_key?(repo_files, safetensors) -> + {:ok, safetensors} + + Map.has_key?(repo_files, pytorch_bin) -> + {:ok, pytorch_bin} + + true -> + {:error, "could not find parameters file in #{dir}"} + end + end + + defp load_tensors(weights_file, weights_path, opts) do + case Path.extname(weights_file) do + ".safetensors" -> + reader = opts[:safetensors_reader] || (&Safetensors.read!(&1, lazy: true)) + reader.(weights_path) + + _ -> + Bumblebee.Conversion.PyTorchLoader.load!(weights_path) + end + end + + defp cast_param(tensor, nil), do: tensor + + defp cast_param(tensor, %Axon.MixedPrecision.Policy{params: type}) do + Nx.as_type(tensor, type) + end + + defp cast_param(tensor, type) do + type = Nx.Type.normalize!(type) + Nx.as_type(tensor, type) + end + + defp allocate_param(tensor, nil), do: tensor + + defp allocate_param(tensor, backend) do + Nx.with_default_backend(backend, fn -> Nx.backend_copy(tensor) end) + end + + defp apply_type(model, nil), do: model + + defp apply_type(model, %Axon.MixedPrecision.Policy{} = policy) do + Axon.MixedPrecision.apply_policy(model, policy) + end + + defp apply_type(model, type) do + type = Nx.Type.normalize!(type) + policy = Axon.MixedPrecision.create_policy(params: type, compute: type, output: type) + Axon.MixedPrecision.apply_policy(model, policy) + end + + defp decode_json(path) do + case File.read(path) do + {:ok, content} -> + case Jason.decode(content) do + {:ok, data} -> {:ok, data} + _ -> {:error, "failed to parse #{path} as JSON"} + end + + {:error, reason} -> + {:error, "failed to read #{path}, reason: #{:file.format_error(reason)}"} + end + end +end diff --git a/lib/bumblebee/text/text_embedding.ex b/lib/bumblebee/text/text_embedding.ex index 0e2f3278..9d123b6c 100644 --- a/lib/bumblebee/text/text_embedding.ex +++ b/lib/bumblebee/text/text_embedding.ex @@ -53,6 +53,9 @@ defmodule Bumblebee.Text.TextEmbedding do %{^output_attribute => output} -> output + %{embedding: output} when output_attribute == :pooled_state -> + output + %{} -> keys = output |> Map.keys() |> Enum.sort() diff --git a/test/bumblebee/huggingface/sentence_transformers_test.exs b/test/bumblebee/huggingface/sentence_transformers_test.exs new file mode 100644 index 00000000..dc5ade06 --- /dev/null +++ b/test/bumblebee/huggingface/sentence_transformers_test.exs @@ -0,0 +1,208 @@ +defmodule Bumblebee.HuggingFace.SentenceTransformersTest do + use ExUnit.Case, async: true + + import Bumblebee.TestHelpers + + @moduletag model_test_tags() + + setup do + model = + Axon.input("input_ids", shape: {nil, nil}) + |> Axon.nx(fn input_ids -> + # Mock base model output with hidden_state {batch_size, seq_len, 4} + batch_size = Nx.axis_size(input_ids, 0) + seq_len = Nx.axis_size(input_ids, 1) + + hidden_state = + Nx.broadcast(1.0, {batch_size, seq_len, 4}) + |> Nx.as_type({:f, 32}) + + %{hidden_state: hidden_state} + end) + + params = Axon.ModelState.empty() + spec = Bumblebee.configure(Bumblebee.Text.Bert, architecture: :base) + + model_info = %{model: model, params: params, spec: spec} + + [model_info: model_info] + end + + describe "load_embedding_head/3" do + @tag :tmp_dir + test "loads pooling and normalize head", %{model_info: model_info, tmp_dir: dir} do + modules = [ + %{"idx" => 0, "path" => "", "type" => "sentence_transformers.models.Transformer"}, + %{"idx" => 1, "path" => "1_Pooling", "type" => "sentence_transformers.models.Pooling"}, + %{"idx" => 2, "path" => "2_Normalize", "type" => "sentence_transformers.models.Normalize"} + ] + + File.write!(Path.join(dir, "modules.json"), Jason.encode!(modules)) + + pooling_dir = Path.join(dir, "1_Pooling") + File.mkdir_p!(pooling_dir) + + pooling_config = %{ + "word_embedding_dimension" => 4, + "pooling_mode_cls_token" => false, + "pooling_mode_mean_tokens" => true, + "pooling_mode_max_tokens" => false + } + + File.write!(Path.join(pooling_dir, "config.json"), Jason.encode!(pooling_config)) + + assert {:ok, new_model_info} = Bumblebee.load_embedding_head({:local, dir}, model_info) + + inputs = %{ + "input_ids" => Nx.tensor([[1, 2, 3]]), + "attention_mask" => Nx.tensor([[1, 1, 0]]) + } + + {_init_fn, predict_fn} = Axon.build(new_model_info.model) + output = predict_fn.(new_model_info.params, inputs) + + assert Map.has_key?(output, :embedding) + assert Map.has_key?(output, :hidden_state) + assert Nx.shape(output.embedding) == {1, 4} + + # Normalization check: L2 norm of output embedding must be 1.0 + norm = Nx.LinAlg.norm(output.embedding, axes: [-1]) + assert_all_close(norm, Nx.tensor([1.0]), atol: 1.0e-5) + end + + @tag :tmp_dir + test "loads dense projection layers with parameters", %{model_info: model_info, tmp_dir: dir} do + modules = [ + %{"idx" => 0, "path" => "", "type" => "sentence_transformers.models.Transformer"}, + %{"idx" => 1, "path" => "1_Pooling", "type" => "sentence_transformers.models.Pooling"}, + %{"idx" => 2, "path" => "2_Dense", "type" => "sentence_transformers.models.Dense"} + ] + + File.write!(Path.join(dir, "modules.json"), Jason.encode!(modules)) + + pooling_dir = Path.join(dir, "1_Pooling") + File.mkdir_p!(pooling_dir) + + File.write!( + Path.join(pooling_dir, "config.json"), + Jason.encode!(%{"pooling_mode_mean_tokens" => true}) + ) + + dense_dir = Path.join(dir, "2_Dense") + File.mkdir_p!(dense_dir) + + dense_config = %{ + "in_features" => 4, + "out_features" => 8, + "bias" => true, + "activation_function" => "torch.nn.modules.linear.Identity" + } + + File.write!(Path.join(dense_dir, "config.json"), Jason.encode!(dense_config)) + + # Weight in PyTorch format: {out_features, in_features} = {8, 4} + weights = %{ + "linear.weight" => Nx.broadcast(0.5, {8, 4}) |> Nx.as_type({:f, 32}), + "linear.bias" => Nx.broadcast(0.1, {8}) |> Nx.as_type({:f, 32}) + } + + Safetensors.write!(Path.join(dense_dir, "model.safetensors"), weights) + + assert {:ok, new_model_info} = Bumblebee.load_embedding_head({:local, dir}, model_info) + + assert Map.has_key?(new_model_info.params.data, "2_Dense") + assert Nx.shape(new_model_info.params.data["2_Dense"]["kernel"]) == {4, 8} + assert Nx.shape(new_model_info.params.data["2_Dense"]["bias"]) == {8} + + inputs = %{ + "input_ids" => Nx.tensor([[1, 2, 3]]), + "attention_mask" => Nx.tensor([[1, 1, 1]]) + } + + {_init_fn, predict_fn} = Axon.build(new_model_info.model) + output = predict_fn.(new_model_info.params, inputs) + + assert Nx.shape(output.embedding) == {1, 8} + end + + @tag :tmp_dir + test "returns error when modules.json is missing", %{model_info: model_info, tmp_dir: dir} do + assert {:error, "could not find modules.json in the repository"} = + Bumblebee.load_embedding_head({:local, dir}, model_info) + end + + @tag :tmp_dir + test "returns error when first module is not Transformer", %{ + model_info: model_info, + tmp_dir: dir + } do + modules = [ + %{"idx" => 0, "path" => "1_Pooling", "type" => "sentence_transformers.models.Pooling"} + ] + + File.write!(Path.join(dir, "modules.json"), Jason.encode!(modules)) + + assert {:error, + "expected the first module in modules.json to be sentence_transformers.models.Transformer"} = + Bumblebee.load_embedding_head({:local, dir}, model_info) + end + + @tag :tmp_dir + test "returns error on unsupported module", %{model_info: model_info, tmp_dir: dir} do + modules = [ + %{"idx" => 0, "path" => "", "type" => "sentence_transformers.models.Transformer"}, + %{"idx" => 1, "path" => "1_Unknown", "type" => "sentence_transformers.models.Unknown"} + ] + + File.write!(Path.join(dir, "modules.json"), Jason.encode!(modules)) + + assert {:error, + "unsupported SentenceTransformers module \"sentence_transformers.models.Unknown\""} = + Bumblebee.load_embedding_head({:local, dir}, model_info) + end + end + + @tag :slow + test "end-to-end with unsloth/embeddinggemma-300m matches reference SentenceTransformers output" do + assert {:ok, model_info} = Bumblebee.load_model({:hf, "unsloth/embeddinggemma-300m"}) + + assert {:ok, model_info} = + Bumblebee.load_embedding_head({:hf, "unsloth/embeddinggemma-300m"}, model_info) + + assert {:ok, tokenizer} = Bumblebee.load_tokenizer({:hf, "unsloth/embeddinggemma-300m"}) + + serving = Bumblebee.Text.text_embedding(model_info, tokenizer) + res = Nx.Serving.run(serving, "Hello world") + + assert Nx.shape(res.embedding) == {768} + + assert_all_close( + res.embedding[0..4], + Nx.tensor([-0.203079, 0.034759, 0.060166, -0.016863, 0.006666]), + atol: 1.0e-4 + ) + end + + @tag :slow + test "end-to-end with sentence-transformers/all-MiniLM-L6-v2" do + assert {:ok, model_info} = + Bumblebee.load_model({:hf, "sentence-transformers/all-MiniLM-L6-v2"}) + + assert {:ok, model_info} = + Bumblebee.load_embedding_head( + {:hf, "sentence-transformers/all-MiniLM-L6-v2"}, + model_info + ) + + assert {:ok, tokenizer} = + Bumblebee.load_tokenizer({:hf, "sentence-transformers/all-MiniLM-L6-v2"}) + + serving = Bumblebee.Text.text_embedding(model_info, tokenizer) + res = Nx.Serving.run(serving, "Hello world") + + assert Nx.shape(res.embedding) == {384} + + norm = Nx.LinAlg.norm(res.embedding) + assert_all_close(norm, Nx.tensor(1.0), atol: 1.0e-5) + end +end