Skip to content
Closed
Show file tree
Hide file tree
Changes from all 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
80 changes: 67 additions & 13 deletions lib/bumblebee.ex
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down
2 changes: 1 addition & 1 deletion lib/bumblebee/huggingface/hub.ex
Original file line number Diff line number Diff line change
Expand Up @@ -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 """
Expand Down
Loading
Loading