Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
Prev Previous commit
Next Next commit
add predict tests
  • Loading branch information
ashleyxuu committed Apr 17, 2024
commit e19f7ac2ce28814c06e21c8fe1082417b4810afc
14 changes: 14 additions & 0 deletions tests/system/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -550,6 +550,20 @@ def llm_fine_tune_df_default_index(
return session.read_gbq(sql)


@pytest.fixture(scope="session")
def llm_remote_text_pandas_df():
"""Additional data matching the penguins dataset, with a new index"""
return pd.DataFrame(
{
"prompt": [
"Please do sentiment analysis on the following text and only output a number from 0 to 5where 0 means sadness, 1 means joy, 2 means love, 3 means anger, 4 means fear, and 5 means surprise. Text: i feel beautifully emotional knowing that these women of whom i knew just a handful were holding me and my baba on our journey",
"Please do sentiment analysis on the following text and only output a number from 0 to 5 where 0 means sadness, 1 means joy, 2 means love, 3 means anger, 4 means fear, and 5 means surprise. Text: i was feeling a little vain when i did this one",
"Please do sentiment analysis on the following text and only output a number from 0 to 5 where 0 means sadness, 1 means joy, 2 means love, 3 means anger, 4 means fear, and 5 means surprise. Text: a father of children killed in an accident",
],
}
)


@pytest.fixture(scope="session")
def time_series_df_default_index(
time_series_table_id: str, session: bigframes.Session
Expand Down
22 changes: 12 additions & 10 deletions tests/system/load/test_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,9 @@
import bigframes.ml.llm


def test_llm_palm_configure_fit(llm_fine_tune_df_default_index, dataset_id):
def test_llm_palm_configure_fit(
llm_fine_tune_df_default_index, llm_remote_text_pandas_df
):
model = bigframes.ml.llm.PaLM2TextGenerator(
model_name="text-bison", max_iterations=1, evaluation_task="CLASSIFICATION"
)
Expand All @@ -25,12 +27,12 @@ def test_llm_palm_configure_fit(llm_fine_tune_df_default_index, dataset_id):
y_train = df[["label"]]
model.fit(X_train, y_train)

# save, load, check parameters to ensure configuration was kept
reloaded_model = model.to_gbq(
f"{dataset_id}.temp_configured_palm_model", replace=True
)
assert (
f"{dataset_id}.temp_configured_palm_model"
in reloaded_model._bqml_model.model_name
)
assert reloaded_model.evaluation_task == "CLASSIFICATION"
assert model is not None

df = model.predict(llm_remote_text_pandas_df).to_pandas()
assert df.shape == (3, 4)
assert "ml_generate_text_llm_result" in df.columns
series = df["ml_generate_text_llm_result"]
assert all(series.str.len() == 1)

# TODO(ashleyxu): After bqml rolled out version control: save, load, check parameters to ensure configuration was kept
Comment thread
ashleyxuu marked this conversation as resolved.
Outdated
12 changes: 0 additions & 12 deletions tests/system/small/ml/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -232,18 +232,6 @@ def palm2_text_generator_model(session, bq_connection) -> llm.PaLM2TextGenerator
return llm.PaLM2TextGenerator(session=session, connection_name=bq_connection)


@pytest.fixture(scope="session")
def palm2_text_generator_fine_tune_model(
session, bq_connection
) -> llm.PaLM2TextGenerator:
return llm.PaLM2TextGenerator(
session=session,
connection_name=bq_connection,
max_iterations=300,
evaluation_task="TEXT_GENERATION",
)


@pytest.fixture(scope="session")
def palm2_text_generator_32k_model(session, bq_connection) -> llm.PaLM2TextGenerator:
return llm.PaLM2TextGenerator(
Expand Down