Skip to content

Commit 97a948c

Browse files
milkshakeiiiashleyxuu
authored andcommitted
feat: support list of numerics in pandas.cut (#580)
An internal user encountered this missing overload
1 parent 9f8f181 commit 97a948c

7 files changed

Lines changed: 227 additions & 1 deletion

File tree

‎bigframes/ml/core.py‎

Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -321,6 +321,46 @@ def create_model(
321321

322322
return self._create_model_with_sql(session=session, sql=sql)
323323

324+
def create_llm_remote_model(
325+
self,
326+
X_train: bpd.DataFrame,
327+
y_train: bpd.DataFrame,
328+
connection_name: str,
329+
options: Mapping[str, Union[str, int, float, Iterable[str]]] = {},
330+
) -> BqmlModel:
331+
"""Create a session-temporary BQML model with the CREATE OR REPLACE MODEL statement
332+
333+
Args:
334+
X_train: features columns for training
335+
y_train: labels columns for training
336+
options: a dict of options to configure the model. Generates a BQML OPTIONS
337+
clause
338+
connection_name:
339+
a BQ connection to talk with Vertex AI, of the format <PROJECT_NUMBER>.<REGION>.<CONNECTION_NAME>. https://cloud.google.com/bigquery/docs/create-cloud-resource-connection
340+
341+
Returns: a BqmlModel, wrapping a trained model in BigQuery
342+
"""
343+
options = dict(options)
344+
# Cache dataframes to make sure base table is not a snapshot
345+
# cached dataframe creates a full copy, never uses snapshot
346+
input_data = X_train._cached(force=True).join(
347+
y_train._cached(force=True), how="outer"
348+
)
349+
options.update({"INPUT_LABEL_COLS": y_train.columns.tolist()})
350+
351+
session = X_train._session
352+
353+
model_ref = self._create_model_ref(session._anonymous_dataset)
354+
355+
sql = self._model_creation_sql_generator.create_llm_remote_model(
356+
source_df=input_data,
357+
model_ref=model_ref,
358+
options=options,
359+
connection_name=connection_name,
360+
)
361+
362+
return self._create_model_with_sql(session=session, sql=sql)
363+
324364
def create_time_series_model(
325365
self,
326366
X_train: bpd.DataFrame,

‎bigframes/ml/llm.py‎

Lines changed: 85 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,11 @@
2727
from bigframes.ml import base, core, globals, utils
2828
import bigframes.pandas as bpd
2929

30+
_BQML_PARAMS_MAPPING = {
31+
"max_iterations": "maxIterations",
32+
"evaluation_task": "evaluationTask",
33+
}
34+
3035
_TEXT_GENERATOR_BISON_ENDPOINT = "text-bison"
3136
_TEXT_GENERATOR_BISON_32K_ENDPOINT = "text-bison-32k"
3237
_TEXT_GENERATOR_ENDPOINTS = (
@@ -51,6 +56,12 @@
5156
class PaLM2TextGenerator(base.BaseEstimator):
5257
"""PaLM2 text generator LLM model.
5358
59+
.. note::
60+
This product or feature is subject to the "Pre-GA Offerings Terms" in the General Service Terms section of the
61+
Service Specific Terms(https://cloud.google.com/terms/service-terms#1). Pre-GA products and features are available "as is"
62+
and might have limited support. For more information, see the launch stage descriptions
63+
(https://cloud.google.com/products#product-launch-stages).
64+
5465
Args:
5566
model_name (str, Default to "text-bison"):
5667
The model for natural language tasks. “text-bison” returns model fine-tuned to follow natural language instructions
@@ -62,6 +73,11 @@ class PaLM2TextGenerator(base.BaseEstimator):
6273
Connection to connect with remote service. str of the format <PROJECT_NUMBER/PROJECT_ID>.<LOCATION>.<CONNECTION_ID>.
6374
if None, use default connection in session context. BigQuery DataFrame will try to create the connection and attach
6475
permission if the connection isn't fully setup.
76+
max_iterations (Optional[int], Default to 300):
77+
The number of steps to run when performing supervised tuning.
78+
evaluation_task (Optional[str], default to "UNSPECIFIED"):
79+
When performing supervised tuning, the type of task that you want to tune the model to perform. Possible values:
80+
"TEXT_GENERATION", "CLASSIFICATION", "SUMMARIZATION", "QUESTION_ANSWERING", "UNSPECIFIED". Default to "UNSPECIFIED".
6581
"""
6682

6783
def __init__(
@@ -70,9 +86,19 @@ def __init__(
7086
model_name: Literal["text-bison", "text-bison-32k"] = "text-bison",
7187
session: Optional[bigframes.Session] = None,
7288
connection_name: Optional[str] = None,
89+
max_iterations: int = 300,
90+
evaluation_task: Literal[
91+
"UNSPECIFIED",
92+
"TEXT_GENERATION",
93+
"CLASSIFICATION",
94+
"SUMMARIZATION",
95+
"QUESTION_ANSWERING",
96+
] = "UNSPECIFIED",
7397
):
7498
self.model_name = model_name
7599
self.session = session or bpd.get_global_session()
100+
self.max_iterations = max_iterations
101+
self.evaluation_task = evaluation_task
76102
self._bq_connection_manager = self.session.bqconnectionmanager
77103

78104
connection_name = connection_name or self.session._bq_connection
@@ -132,12 +158,70 @@ def _from_bq(
132158
model_connection = model._properties["remoteModelInfo"]["connection"]
133159
model_endpoint = bqml_endpoint.split("/")[-1]
134160

161+
# Get the optional params
162+
kwargs: dict = {}
163+
last_fitting = model.training_runs[-1]["trainingOptions"]
164+
165+
dummy_arima = cls()
166+
for bf_param, _ in dummy_arima.__dict__.items():
167+
bqml_param = _BQML_PARAMS_MAPPING.get(bf_param)
168+
if bqml_param in last_fitting:
169+
# Convert types
170+
if bf_param in ["max_iterations"]:
171+
kwargs[bf_param] = int(last_fitting[bqml_param])
172+
elif bf_param in ["evaluation_task"]:
173+
kwargs[bf_param] = str(last_fitting[bqml_param])
174+
135175
text_generator_model = cls(
136-
session=session, model_name=model_endpoint, connection_name=model_connection
176+
**kwargs,
177+
session=session,
178+
model_name=model_endpoint,
179+
connection_name=model_connection,
137180
)
138181
text_generator_model._bqml_model = core.BqmlModel(session, model)
139182
return text_generator_model
140183

184+
@property
185+
def _bqml_options(self) -> dict:
186+
"""The model options as they will be set for BQML"""
187+
options = {
188+
"max_iterations": self.max_iterations,
189+
"data_split_method": "NO_SPLIT",
190+
"evaluation_task": self.evaluation_task,
191+
}
192+
return options
193+
194+
def fit(
195+
self,
196+
X: Union[bpd.DataFrame, bpd.Series],
197+
y: Union[bpd.DataFrame, bpd.Series],
198+
) -> PaLM2TextGenerator:
199+
"""Fine tune PaLM2TextGenerator model.
200+
201+
Args:
202+
X (bigframes.dataframe.DataFrame or bigframes.series.Series):
203+
DataFrame of shape (n_samples, n_features). Training data.
204+
y (bigframes.dataframe.DataFrame or bigframes.series.Series:
205+
Training labels.
206+
207+
Returns:
208+
PaLM2TextGenerator: Fitted Estimator.
209+
"""
210+
X, y = utils.convert_to_dataframe(X, y)
211+
212+
# TODO(ashleyxu): options= self._bqml_options
213+
options = self._bqml_options
214+
options["endpoint"] = self.model_name + "@001"
215+
options["prompt_col"] = X.columns.tolist()[0]
216+
217+
self._bqml_model = self._bqml_model_factory.create_llm_remote_model(
218+
X,
219+
y,
220+
options=options,
221+
connection_name=self.connection_name,
222+
)
223+
return self
224+
141225
def predict(
142226
self,
143227
X: Union[bpd.DataFrame, bpd.Series],

‎bigframes/ml/sql.py‎

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -177,6 +177,24 @@ def create_model(
177177
parts.append(f"AS {source_sql}")
178178
return "\n".join(parts)
179179

180+
# Model create and alter
181+
def create_llm_remote_model(
182+
self,
183+
source_df: bpd.DataFrame,
184+
connection_name: str,
185+
model_ref: google.cloud.bigquery.ModelReference,
186+
options: Mapping[str, Union[str, int, float, Iterable[str]]] = {},
187+
) -> str:
188+
"""Encode the CREATE OR REPLACE MODEL statement for BQML"""
189+
source_sql = source_df.sql
190+
191+
parts = [f"CREATE OR REPLACE MODEL {self._model_id_sql(model_ref)}"]
192+
parts.append(self.connection(connection_name))
193+
if options:
194+
parts.append(self.options(**options))
195+
parts.append(f"AS {source_sql}")
196+
return "\n".join(parts)
197+
180198
def create_remote_model(
181199
self,
182200
connection_name: str,

‎tests/system/conftest.py‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -537,6 +537,19 @@ def penguins_df_default_index(
537537
return session.read_gbq(penguins_table_id)
538538

539539

540+
@pytest.fixture(scope="session")
541+
def llm_fine_tune_df_default_index(
542+
session: bigframes.Session,
543+
) -> bigframes.dataframe.DataFrame:
544+
sql = """
545+
SELECT
546+
CONCAT("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: ", text) as prompt,
547+
CAST(label AS STRING) as label
548+
FROM `llm_tuning.emotion_classification_train`
549+
"""
550+
return session.read_gbq(sql)
551+
552+
540553
@pytest.fixture(scope="session")
541554
def time_series_df_default_index(
542555
time_series_table_id: str, session: bigframes.Session

‎tests/system/large/ml/test_llm.py‎

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,36 @@
1+
# Copyright 2023 Google LLC
2+
#
3+
# Licensed under the Apache License, Version 2.0 (the "License");
4+
# you may not use this file except in compliance with the License.
5+
# You may obtain a copy of the License at
6+
#
7+
# http://www.apache.org/licenses/LICENSE-2.0
8+
#
9+
# Unless required by applicable law or agreed to in writing, software
10+
# distributed under the License is distributed on an "AS IS" BASIS,
11+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
# See the License for the specific language governing permissions and
13+
# limitations under the License.
14+
15+
import bigframes.ml.llm
16+
17+
18+
def test_llm_palm_configure_fit(llm_fine_tune_df_default_index, dataset_id):
19+
model = bigframes.ml.llm.PaLM2TextGenerator(
20+
model_name="text-bison", max_iterations=1, evaluation_task="CLASSIFICATION"
21+
)
22+
23+
df = llm_fine_tune_df_default_index.dropna()
24+
X_train = df[["prompt"]]
25+
y_train = df[["label"]]
26+
model.fit(X_train, y_train)
27+
28+
# save, load, check parameters to ensure configuration was kept
29+
reloaded_model = model.to_gbq(
30+
f"{dataset_id}.temp_configured_palm_model", replace=True
31+
)
32+
assert (
33+
f"{dataset_id}.temp_configured_palm_model"
34+
in reloaded_model._bqml_model.model_name
35+
)
36+
assert reloaded_model.evaluation_task == "CLASSIFICATION"

‎tests/system/small/ml/conftest.py‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -232,6 +232,18 @@ def palm2_text_generator_model(session, bq_connection) -> llm.PaLM2TextGenerator
232232
return llm.PaLM2TextGenerator(session=session, connection_name=bq_connection)
233233

234234

235+
@pytest.fixture(scope="session")
236+
def palm2_text_generator_fine_tune_model(
237+
session, bq_connection
238+
) -> llm.PaLM2TextGenerator:
239+
return llm.PaLM2TextGenerator(
240+
session=session,
241+
connection_name=bq_connection,
242+
max_iterations=300,
243+
evaluation_task="TEXT_GENERATION",
244+
)
245+
246+
235247
@pytest.fixture(scope="session")
236248
def palm2_text_generator_32k_model(session, bq_connection) -> llm.PaLM2TextGenerator:
237249
return llm.PaLM2TextGenerator(

‎tests/unit/ml/test_sql.py‎

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -181,6 +181,29 @@ def test_create_model_transform_correct(
181181
)
182182

183183

184+
def test_create_llm_remote_model_correct(
185+
model_creation_sql_generator: ml_sql.ModelCreationSqlGenerator,
186+
mock_df: bpd.DataFrame,
187+
):
188+
sql = model_creation_sql_generator.create_llm_remote_model(
189+
source_df=mock_df,
190+
connection_name="my_project.us.my_connection",
191+
model_ref=bigquery.ModelReference.from_string(
192+
"test-proj._anonXYZ.create_remote_model"
193+
),
194+
options={"option_key1": "option_value1", "option_key2": 2},
195+
)
196+
assert (
197+
sql
198+
== """CREATE OR REPLACE MODEL `test-proj`.`_anonXYZ`.`create_remote_model`
199+
REMOTE WITH CONNECTION `my_project.us.my_connection`
200+
OPTIONS(
201+
option_key1="option_value1",
202+
option_key2=2)
203+
AS input_X_y_sql"""
204+
)
205+
206+
184207
def test_create_remote_model_correct(
185208
model_creation_sql_generator: ml_sql.ModelCreationSqlGenerator,
186209
):

0 commit comments

Comments
 (0)