Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@

from bigframes import series, session
from bigframes.core import col, sentinels
from bigframes.extensions.core import abstract_series_accessor, series_tvf_mixins
from bigframes.extensions.core import abstract_series_accessor, series_mixins

T = TypeVar("T")
S = TypeVar("S")
Expand Down Expand Up @@ -1112,7 +1112,7 @@ def unix_date(
return self._to_series(cast(series.Series, result))


class AiSeriesAccessor(series_tvf_mixins.AITVFMixin[T, S]):
class AiSeriesAccessor(series_mixins.AIMixin[T, S]):
"""Series accessor for BigQuery ai functions."""


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -14,19 +14,21 @@

from __future__ import annotations

from typing import List, Mapping, TypeVar
from typing import Any, List, Literal, Mapping, TypeVar

import pandas as pd

from bigframes import session
from bigframes import series
from bigframes import session as bf_session
from bigframes.bigquery import ai
from bigframes.extensions.core import abstract_series_accessor
from bigframes.ml import base as ml_base

T = TypeVar("T")
S = TypeVar("S")


class AITVFMixin(abstract_series_accessor.AbstractBigQuerySeriesAccessor[T, S]):
class AIMixin(abstract_series_accessor.AbstractBigQuerySeriesAccessor[T, S]):
def generate_embedding(
self,
model: ml_base.BaseEstimator | str | pd.Series,
Expand All @@ -37,18 +39,17 @@ def generate_embedding(
end_second: float | None = None,
interval_seconds: float | None = None,
trial_id: int | None = None,
session: session.Session | None = None,
session: bf_session.Session | None = None,
) -> T:
"""
Creates embeddings that describe an entity — for example, a piece of text or an image.

This is an accessor for :func:`bigframes.bigquery.ai.generate_embedding`. See that
function's documentation for detailed parameter descriptions and examples.
"""
import bigframes.bigquery.ai

bf_series = self._bf_from_series(session)
result = bigframes.bigquery.ai.generate_embedding(
result = ai.generate_embedding(
model,
bf_series,
output_dimensionality=output_dimensionality,
Expand All @@ -71,18 +72,16 @@ def generate_text(
stop_sequences: List[str] | None = None,
ground_with_google_search: bool | None = None,
request_type: str | None = None,
session: session.Session | None = None,
session: bf_session.Session | None = None,
) -> T:
"""
Generates text using a BigQuery ML model.

This is an accessor for :func:`bigframes.bigquery.ai.generate_text`. See that
function's documentation for detailed parameter descriptions and examples.
"""
import bigframes.bigquery.ai

bf_series = self._bf_from_series(session)
result = bigframes.bigquery.ai.generate_text(
result = ai.generate_text(
model,
bf_series,
temperature=temperature,
Expand All @@ -105,18 +104,16 @@ def generate_table(
max_output_tokens: int | None = None,
stop_sequences: List[str] | None = None,
request_type: str | None = None,
session: session.Session | None = None,
session: bf_session.Session | None = None,
) -> T:
"""
Generates a table using a BigQuery ML model.

This is an accessor for :func:`bigframes.bigquery.ai.generate_table`. See that
function's documentation for detailed parameter descriptions and examples.
"""
import bigframes.bigquery.ai

bf_series = self._bf_from_series(session)
result = bigframes.bigquery.ai.generate_table(
result = ai.generate_table(
model,
bf_series,
output_schema=output_schema,
Expand All @@ -127,3 +124,73 @@ def generate_table(
request_type=request_type,
)
return self._to_dataframe(result)

def embed(
self,
*,
endpoint: str | None = None,
model: str | None = None,
task_type: (
Literal[
"retrieval_query",
"retrieval_document",
"semantic_similarity",
"classification",
"clustering",
"question_answering",
"fact_verification",
"code_retrieval_query",
]
| None
) = None,
title: str | None = None,
model_params: Mapping[Any, Any] | None = None,
connection_id: str | None = None,
session: bf_session.Session | None = None,
) -> S:
"""
Creates embeddings from text or image data in BigQuery.

This is an accessor for :func:`bigframes.bigquery.ai.embed`. See that
function's documentation for detailed parameter descriptions and examples.
"""

bf_series = self._bf_from_series(session)
result = ai.embed(
bf_series,
endpoint=endpoint,
model=model,
task_type=task_type,
title=title,
model_params=model_params,
connection_id=connection_id,
)
return self._to_series(result)
Comment thread
sycai marked this conversation as resolved.

def similarity(
self,
other: str | series.Series | pd.Series,
*,
endpoint: str | None = None,
model: str | None = None,
model_params: Mapping[Any, Any] | None = None,
connection_id: str | None = None,
session: bf_session.Session | None = None,
) -> S:
"""
Returns a FLOAT64 value that represents the cosine similarity between the two inputs.

This is an accessor for :func:`bigframes.bigquery.ai.similarity`. See that
function's documentation for detailed parameter descriptions and examples.
"""

bf_series = self._bf_from_series(session)
result = ai.similarity(
bf_series,
other,
endpoint=endpoint,
model=model,
model_params=model_params,
connection_id=connection_id,
)
return self._to_series(result)
Comment thread
sycai marked this conversation as resolved.
Original file line number Diff line number Diff line change
Expand Up @@ -20,15 +20,15 @@ from typing import (

from bigframes import series, session
from bigframes.core import col, sentinels
from bigframes.extensions.core import abstract_series_accessor, series_tvf_mixins
from bigframes.extensions.core import abstract_series_accessor, series_mixins

T = TypeVar("T")
S = TypeVar("S")


{% for ns in namespaces %}
{% if ns.class_name == "AiSeriesAccessor" %}
class {{ ns.class_name }}(series_tvf_mixins.AITVFMixin[T, S]):
class {{ ns.class_name }}(series_mixins.AIMixin[T, S]):
{% else %}
class {{ ns.class_name }}(abstract_series_accessor.AbstractBigQuerySeriesAccessor[T, S]):
{% endif %}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -246,3 +246,139 @@ def test_bigframes_ai_generate_table(scalar_types_df: bpd.DataFrame, monkeypatch
}
result_df.to_pandas.assert_not_called()
assert actual_result is result_df


def test_ai_embed(monkeypatch):
session = mock.create_autospec(bigframes.session.Session)
bf_series = mock.create_autospec(bpd.Series)
session.read_pandas.return_value = bf_series

mock_embed = mock.MagicMock()
result_series = mock.create_autospec(bpd.Series)
mock_embed.return_value = result_series
expected_result = mock.create_autospec(pd.Series)
result_series.to_pandas.return_value = expected_result

monkeypatch.setattr(bigframes.bigquery.ai, "embed", mock_embed)

series = pd.Series(["hello world"], name="content")
actual_result = series.bigquery.ai.embed( # type: ignore
endpoint="my_endpoint",
model="my_model",
task_type="retrieval_query",
title="my_title",
model_params={"key": "val"},
connection_id="my_connection",
session=session,
)

session.read_pandas.assert_called_once()
mock_embed.assert_called_once_with(
bf_series,
endpoint="my_endpoint",
model="my_model",
task_type="retrieval_query",
title="my_title",
model_params={"key": "val"},
connection_id="my_connection",
)
result_series.to_pandas.assert_called_once()
assert actual_result is expected_result


def test_bigframes_ai_embed(scalar_types_df: bpd.DataFrame, monkeypatch):
session = mock.create_autospec(bigframes.session.Session)
result_series = mock.create_autospec(bpd.Series)

mock_embed = mock.MagicMock()
mock_embed.return_value = result_series

monkeypatch.setattr(bigframes.bigquery.ai, "embed", mock_embed)

scalar_types_series = scalar_types_df["string_col"]
actual_result = scalar_types_series.bigquery.ai.embed(
endpoint="my_endpoint",
session=session,
)

session.read_pandas.assert_not_called()
mock_embed.assert_called_once()
args, kwargs = mock_embed.call_args
assert args[0] is scalar_types_series
assert kwargs == {
"endpoint": "my_endpoint",
"model": None,
"task_type": None,
"title": None,
"model_params": None,
"connection_id": None,
}
result_series.to_pandas.assert_not_called()
assert actual_result is result_series


def test_ai_similarity(monkeypatch):
session = mock.create_autospec(bigframes.session.Session)
bf_series = mock.create_autospec(bpd.Series)
session.read_pandas.return_value = bf_series

mock_similarity = mock.MagicMock()
result_series = mock.create_autospec(bpd.Series)
mock_similarity.return_value = result_series
expected_result = mock.create_autospec(pd.Series)
result_series.to_pandas.return_value = expected_result

monkeypatch.setattr(bigframes.bigquery.ai, "similarity", mock_similarity)

series = pd.Series(["apple"], name="content")
actual_result = series.bigquery.ai.similarity( # type: ignore
"banana",
endpoint="my_endpoint",
model="my_model",
model_params={"key": "val"},
connection_id="my_connection",
session=session,
)

session.read_pandas.assert_called_once()
mock_similarity.assert_called_once_with(
bf_series,
"banana",
endpoint="my_endpoint",
model="my_model",
model_params={"key": "val"},
connection_id="my_connection",
)
result_series.to_pandas.assert_called_once()
assert actual_result is expected_result


def test_bigframes_ai_similarity(scalar_types_df: bpd.DataFrame, monkeypatch):
session = mock.create_autospec(bigframes.session.Session)
result_series = mock.create_autospec(bpd.Series)

mock_similarity = mock.MagicMock()
mock_similarity.return_value = result_series

monkeypatch.setattr(bigframes.bigquery.ai, "similarity", mock_similarity)

scalar_types_series = scalar_types_df["string_col"]
actual_result = scalar_types_series.bigquery.ai.similarity(
"other_text",
endpoint="my_endpoint",
session=session,
)

session.read_pandas.assert_not_called()
mock_similarity.assert_called_once()
args, kwargs = mock_similarity.call_args
assert args[0] is scalar_types_series
assert args[1] == "other_text"
assert kwargs == {
"endpoint": "my_endpoint",
"model": None,
"model_params": None,
"connection_id": None,
}
result_series.to_pandas.assert_not_called()
assert actual_result is result_series
Loading