ML en Snowflake a escala con Snowpark Python y XGBoost
Durante el último año y medio, nuestro equipo ha estado trabajando con cientos de clientes que desarrollan Snowflake con Snowpark y Python (por cierto, consulte esta publicación de mi colega Caleb Baechtold que recopila todo lo aprendido sobre la puesta en funcionamiento de Snowpark Python en producción) y un La pregunta que recibo con frecuencia es: "pero tenemos datos realmente grandes, ¿cómo funcionan a escala?".
Muchos de los ejemplos/demostraciones que encontrará en línea son excelentes para mostrar ejemplos artificiales simples que muestran la funcionalidad principal, pero pocos, si es que hay alguno, muestran ejemplos de cómo funciona en el contexto de problemas a gran escala empresarial.
TPC-DS , es un conjunto de datos útil a gran escala que generalmente representa un esquema de datos comerciales genéricos (no estoy aquí para defender que TPC-DS, o cualquier conjunto de datos sea el final, todo sea un punto de referencia de cualquier plataforma o herramienta, solo creo que puede ser un conjunto de datos útil como punto de partida para probar cosas a escala). Snowflake pone a disposición TPC-DS en cada cuenta de Snowflake como un recurso compartido de datos, tanto en una edición de 10 TB como de 100 TB (por cierto, porque está expuesto como un recurso compartido de datos, usted como usuario no paga ninguno de los costos de almacenamiento . Solo para cualquier cómputo potencial que use para consultarlo). La edición de 100 TB tiene más de 560 mil millones de filas en las tablas de hechos. La edición de 10 TB, 56 mil millones.
Podemos usar Snowpark Python para crear una solución de ML relativamente simple para un problema comercial común que tiene un negocio hipotético como TPC que es "Quiero poder predecir el valor de por vida de mis clientes en todos los canales de ventas". Usaremos la API de marco de datos de Python de Snowpark para realizar la preparación de datos/ingeniería de características, los procedimientos almacenados y los almacenes optimizados de Snowpark para el entrenamiento y los UDF por lotes para la inferencia, todo sin que los datos tengan que salir de Snowflake al utilizar los recursos informáticos y la capacidad de escala de Snowflake.
Comencemos con nuestro código de preparación de datos/ingeniería de características. Bastante simple, agregaremos las ventas por cliente a través de todos los canales. Luego lo uniremos a las tablas de dimensiones de nuestros clientes para obtener características potenciales de interés.
store_sales_agged = store_sales.group_by('ss_customer_sk').agg(F.sum('ss_sales_price').as_('total_sales'))
web_sales_agged = web_sales.group_by('ws_bill_customer_sk').agg(F.sum('ws_sales_price').as_('total_sales'))
catalog_sales_agged = catalog_sales.group_by('cs_bill_customer_sk').agg(F.sum('cs_sales_price').as_('total_sales'))
store_sales_agged = store_sales_agged.rename('ss_customer_sk', 'customer_sk')
web_sales_agged = web_sales_agged.rename('ws_bill_customer_sk', 'customer_sk')
catalog_sales_agged = catalog_sales_agged.rename('cs_bill_customer_sk', 'customer_sk')
total_sales = store_sales_agged.union_all(web_sales_agged)
total_sales = total_sales.union_all(catalog_sales_agged)
total_sales = total_sales.group_by('customer_sk').agg(F.sum('total_sales').as_('total_sales'))
customer = customer.select('c_customer_sk','c_current_hdemo_sk', 'c_current_addr_sk', 'c_customer_id', 'c_birth_year')
customer = customer.join(address.select('ca_address_sk', 'ca_zip'), customer['c_current_addr_sk'] == address['ca_address_sk'] )
customer = customer.join(demo.select('cd_demo_sk', 'cd_gender', 'cd_marital_status', 'cd_credit_rating', 'cd_education_status', 'cd_dep_count'),
customer['c_current_hdemo_sk'] == demo['cd_demo_sk'] )
customer = customer.rename('c_customer_sk', 'customer_sk')
final_df = total_sales.join(customer, on='customer_sk')
session.use_database('tpcds_xgboost')
session.use_schema('demo')
final_df.write.mode('overwrite').save_as_table('feature_store')
Ahora estamos listos para entrenar nuestro modelo usando un procedimiento almacenado. Snowflake simplemente pone a su disposición todos los marcos populares de Python ML para que los utilice para la capacitación de ML a través de nuestra asociación con Anaconda . No es necesario aprender una nueva biblioteca o sintaxis.
from sklearn.pipeline import Pipeline
from sklearn.impute import SimpleImputer
from sklearn.preprocessing import StandardScaler, OneHotEncoder, MinMaxScaler
from sklearn.metrics import mean_squared_error
from sklearn.compose import ColumnTransformer
from xgboost import XGBRegressor
import joblib
import os
def train_model(session: snowflake.snowpark.Session) -> float:
snowdf = session.table("feature_store")
snowdf = snowdf.drop(['CUSTOMER_SK', 'C_CURRENT_HDEMO_SK', 'C_CURRENT_ADDR_SK', 'C_CUSTOMER_ID', 'CA_ADDRESS_SK', 'CD_DEMO_SK'])
snowdf_train, snowdf_test = snowdf.random_split([0.8, 0.2], seed=82)
# save the train and test sets as time stamped tables in Snowflake
snowdf_train.write.mode("overwrite").save_as_table("tpcds_xgboost.demo.tpc_TRAIN")
snowdf_test.write.mode("overwrite").save_as_table("tpcds_xgboost.demo.tpc_TEST")
train_x = snowdf_train.drop("TOTAL_SALES").to_pandas() # drop labels for training set
train_y = snowdf_train.select("TOTAL_SALES").to_pandas()
test_x = snowdf_test.drop("TOTAL_SALES").to_pandas()
test_y = snowdf_test.select("TOTAL_SALES").to_pandas()
cat_cols = ['CA_ZIP', 'CD_GENDER', 'CD_MARITAL_STATUS', 'CD_CREDIT_RATING', 'CD_EDUCATION_STATUS']
num_cols = ['C_BIRTH_YEAR', 'CD_DEP_COUNT']
num_pipeline = Pipeline([
('imputer', SimpleImputer(strategy="median")),
('std_scaler', StandardScaler()),
])
preprocessor = ColumnTransformer(
transformers=[('num', num_pipeline, num_cols),
('encoder', OneHotEncoder(handle_unknown="ignore"), cat_cols) ])
pipe = Pipeline([('preprocessor', preprocessor),
('xgboost', XGBRegressor())])
pipe.fit(train_x, train_y)
test_preds = pipe.predict(test_x)
rmse = mean_squared_error(test_y, test_preds)
model_file = os.path.join('/tmp', 'model.joblib')
joblib.dump(pipe, model_file)
session.file.put(model_file, "@ml_models",overwrite=True)
return rmse
train_model_sp = F.sproc(train_model, session=session, replace=True)
# Switch to Snowpark Optimized Warehouse for training and to run the stored proc
session.use_warehouse('snowpark_opt_wh')
train_model_sp(session=session)
import sys
import pandas as pd
import cachetools
import joblib
from snowflake.snowpark import types as T
session.add_import("@ml_models/model.joblib")
features = [ 'C_BIRTH_YEAR', 'CA_ZIP', 'CD_GENDER', 'CD_MARITAL_STATUS', 'CD_CREDIT_RATING', 'CD_EDUCATION_STATUS', 'CD_DEP_COUNT']
@cachetools.cached(cache={})
def read_file(filename):
import_dir = sys._xoptions.get("snowflake_import_directory")
if import_dir:
with open(os.path.join(import_dir, filename), 'rb') as file:
m = joblib.load(file)
return m
@F.pandas_udf(session=session, max_batch_size=10000, is_permanent=True,
stage_location='@ml_models', name="clv_xgboost_udf")
def predict(df: T.PandasDataFrame[int, str, str, str, str, str, int]) -> T.PandasSeries[float]:
m = read_file('model.joblib')
df.columns = features
return m.predict(df)
inference_df = session.table('feature_store')
inference_df = inference_df.drop(['CUSTOMER_SK', 'C_CURRENT_HDEMO_SK', 'C_CURRENT_ADDR_SK', 'C_CUSTOMER_ID', 'CA_ADDRESS_SK', 'CD_DEMO_SK'])
inputs = inference_df.drop("TOTAL_SALES")
snowdf_results = inference_df.select(*inputs,
predict(*inputs).alias('PREDICTION'),
(F.col('TOTAL_SALES')).alias('ACTUAL_SALES')
)
snowdf_results.write.mode('overwrite').save_as_table('predictions')
SELECT "C_BIRTH_YEAR",
"CA_ZIP",
"CD_GENDER",
"CD_MARITAL_STATUS",
"CD_CREDIT_RATING",
"CD_EDUCATION_STATUS",
"CD_DEP_COUNT",
clv_xgboost_udf("C_BIRTH_YEAR", "CA_ZIP", "CD_GENDER", "CD_MARITAL_STATUS", "CD_CREDIT_RATING", "CD_EDUCATION_STATUS", "CD_DEP_COUNT") AS "PREDICTION",
"TOTAL_SALES" AS "ACTUAL_SALES"
FROM tpcds_xgboost.demo.feature_store
Como les digo a todos mis clientes con los que trabajo, por favor no confíen en mi palabra. Pruébelo usted mismo, todo el código está disponible para usted aquí . Incluso mejor que TPC-DS, pruébelo con algunos de los datos y canalizaciones de su organización.

![¿Qué es una lista vinculada, de todos modos? [Parte 1]](https://post.nghiatu.com/assets/images/m/max/724/1*Xokk6XOjWyIGCBujkJsCzQ.jpeg)



































