Use sqlite vector
This commit is contained in:
@@ -10,7 +10,7 @@ defmodule MusicLibrary.Records.RecordEmbedding do
|
||||
schema "record_embeddings" do
|
||||
belongs_to :record, Record
|
||||
|
||||
field :embedding, MusicLibrary.Records.RecordEmbedding.EmbeddingType
|
||||
field :embedding, SqliteVec.Ecto.Float32
|
||||
field :text_representation, :string
|
||||
|
||||
timestamps(type: :utc_datetime)
|
||||
@@ -20,28 +20,6 @@ defmodule MusicLibrary.Records.RecordEmbedding do
|
||||
record_embedding
|
||||
|> cast(attrs, [:record_id, :embedding, :text_representation])
|
||||
|> validate_required([:record_id, :embedding, :text_representation])
|
||||
|> validate_embedding_dimensions()
|
||||
|> unique_constraint(:record_id)
|
||||
end
|
||||
|
||||
defp validate_embedding_dimensions(changeset) do
|
||||
case get_change(changeset, :embedding) do
|
||||
nil ->
|
||||
changeset
|
||||
|
||||
embedding when is_list(embedding) ->
|
||||
if length(embedding) == 1536 do
|
||||
changeset
|
||||
else
|
||||
add_error(
|
||||
changeset,
|
||||
:embedding,
|
||||
"must have exactly 1536 dimensions, got #{length(embedding)}"
|
||||
)
|
||||
end
|
||||
|
||||
_ ->
|
||||
add_error(changeset, :embedding, "must be a list of floats")
|
||||
end
|
||||
end
|
||||
end
|
||||
|
||||
@@ -1,50 +0,0 @@
|
||||
defmodule MusicLibrary.Records.RecordEmbedding.EmbeddingType do
|
||||
@moduledoc """
|
||||
Custom Ecto type for storing embedding vectors.
|
||||
|
||||
Embeddings are stored as JSON-encoded arrays of floats in the database,
|
||||
but presented as Elixir lists in the application.
|
||||
"""
|
||||
use Ecto.Type
|
||||
|
||||
@impl true
|
||||
def type, do: :string
|
||||
|
||||
@impl true
|
||||
def cast(embedding) when is_list(embedding) do
|
||||
if Enum.all?(embedding, &is_float/1) or Enum.all?(embedding, &is_number/1) do
|
||||
# Convert all numbers to floats
|
||||
{:ok, Enum.map(embedding, &to_float/1)}
|
||||
else
|
||||
:error
|
||||
end
|
||||
end
|
||||
|
||||
def cast(_), do: :error
|
||||
|
||||
@impl true
|
||||
def load(json) when is_binary(json) do
|
||||
case JSON.decode(json) do
|
||||
{:ok, embedding} when is_list(embedding) ->
|
||||
{:ok, Enum.map(embedding, &to_float/1)}
|
||||
|
||||
_ ->
|
||||
:error
|
||||
end
|
||||
end
|
||||
|
||||
def load(_), do: :error
|
||||
|
||||
@impl true
|
||||
def dump(embedding) when is_list(embedding) do
|
||||
json = JSON.encode!(embedding)
|
||||
{:ok, json}
|
||||
rescue
|
||||
_ -> :error
|
||||
end
|
||||
|
||||
def dump(_), do: :error
|
||||
|
||||
defp to_float(n) when is_float(n), do: n
|
||||
defp to_float(n) when is_integer(n), do: n * 1.0
|
||||
end
|
||||
@@ -4,7 +4,9 @@ defmodule MusicLibrary.Records.Similarity do
|
||||
"""
|
||||
|
||||
import Ecto.Query
|
||||
import(SqliteVec.Ecto.Query)
|
||||
|
||||
alias MusicLibrary.Records
|
||||
alias MusicLibrary.Records.{Record, RecordEmbedding}
|
||||
alias MusicLibrary.Repo
|
||||
|
||||
@@ -53,44 +55,30 @@ defmodule MusicLibrary.Records.Similarity do
|
||||
"""
|
||||
def find_similar(record_id, opts \\ []) do
|
||||
limit = Keyword.get(opts, :limit, 10)
|
||||
min_similarity = Keyword.get(opts, :min_similarity, 0.0)
|
||||
scope = Keyword.get(opts, :scope)
|
||||
|
||||
with {:ok, source_embedding} <- get_embedding(record_id),
|
||||
similar_records <- calculate_similarities(source_embedding, record_id, scope) do
|
||||
similar_records
|
||||
|> Enum.filter(fn {_record, similarity} -> similarity >= min_similarity end)
|
||||
|> Enum.take(limit)
|
||||
|> Enum.map(fn {record, similarity} -> {record, Float.round(similarity, 4)} end)
|
||||
else
|
||||
{:error, :not_found} -> []
|
||||
end
|
||||
end
|
||||
record = Records.get_record!(record_id)
|
||||
record_musicbrainz_id = record.musicbrainz_id
|
||||
|
||||
@doc """
|
||||
Calculates cosine similarity between two embedding vectors.
|
||||
case get_embedding(record_id) do
|
||||
{:ok, source_embedding} ->
|
||||
query =
|
||||
from re in RecordEmbedding,
|
||||
where: re.record_id != ^record_id,
|
||||
join: r in Record,
|
||||
on: r.id == re.record_id and r.musicbrainz_id != ^record_musicbrainz_id,
|
||||
order_by: vec_distance_cosine(re.embedding, vec_f32(source_embedding)),
|
||||
select: {r, re.embedding},
|
||||
group_by: r.musicbrainz_id,
|
||||
limit: ^limit
|
||||
|
||||
Returns a float between -1.0 and 1.0, where:
|
||||
- 1.0 = identical vectors
|
||||
- 0.0 = orthogonal vectors
|
||||
- -1.0 = opposite vectors
|
||||
"""
|
||||
def cosine_similarity(vec_a, vec_b) when is_list(vec_a) and is_list(vec_b) do
|
||||
if length(vec_a) != length(vec_b) do
|
||||
raise ArgumentError, "Vectors must have the same length"
|
||||
end
|
||||
query = apply_scope_filter(query, scope)
|
||||
|
||||
dot_product =
|
||||
Enum.zip(vec_a, vec_b)
|
||||
|> Enum.reduce(0.0, fn {a, b}, acc -> acc + a * b end)
|
||||
query
|
||||
|> Repo.all()
|
||||
|
||||
magnitude_a = calculate_magnitude(vec_a)
|
||||
magnitude_b = calculate_magnitude(vec_b)
|
||||
|
||||
if magnitude_a == 0.0 or magnitude_b == 0.0 do
|
||||
0.0
|
||||
else
|
||||
dot_product / (magnitude_a * magnitude_b)
|
||||
{:error, :not_found} ->
|
||||
[]
|
||||
end
|
||||
end
|
||||
|
||||
@@ -142,31 +130,6 @@ defmodule MusicLibrary.Records.Similarity do
|
||||
defp humanize_type(:other), do: "Other"
|
||||
defp humanize_type(_), do: "Unknown"
|
||||
|
||||
defp calculate_magnitude(vector) do
|
||||
vector
|
||||
|> Enum.reduce(0.0, fn x, acc -> acc + x * x end)
|
||||
|> :math.sqrt()
|
||||
end
|
||||
|
||||
defp calculate_similarities(source_embedding, source_record_id, scope) do
|
||||
query =
|
||||
from re in RecordEmbedding,
|
||||
where: re.record_id != ^source_record_id,
|
||||
join: r in Record,
|
||||
on: r.id == re.record_id,
|
||||
select: {r, re.embedding}
|
||||
|
||||
query = apply_scope_filter(query, scope)
|
||||
|
||||
query
|
||||
|> Repo.all()
|
||||
|> Enum.map(fn {record, embedding} ->
|
||||
similarity = cosine_similarity(source_embedding, embedding)
|
||||
{record, similarity}
|
||||
end)
|
||||
|> Enum.sort_by(fn {_record, similarity} -> similarity end, :desc)
|
||||
end
|
||||
|
||||
defp apply_scope_filter(query, :collection) do
|
||||
from [re, r] in query, where: not is_nil(r.purchased_at)
|
||||
end
|
||||
|
||||
Reference in New Issue
Block a user