如何将Snowflake ARRAY列作为Snowpark分解模型的输入?
Snowpark中处理ARRAY列的TruncatedSVD降维方案
问题原因
Snowpark ML的TruncatedSVD封装了sklearn的实现,它要求输入为多个数值特征列,而非单个ARRAY列。直接传入ARRAY列时,底层会将其当作列名列表解析,自然触发KeyError。
解决方案一:固定长度数组转多数值列(适合小尺寸数组)
如果嵌入数组长度固定,可先将ARRAY列展开为多个独立数值列,再传入TruncatedSVD:
# 1. 展开ARRAY列为多数值列(示例数组长度为4) vec_length = 4 expand_cols = ", ".join([f"doc_vec[{i}] as vec_{i}" for i in range(vec_length)]) df_expanded = session.sql(f""" select doc_id, {expand_cols} from ( select 'doc1' as doc_id, array_construct(0.1, 0.3, 0.5, 0.7) as doc_vec union select 'doc2' as doc_id, array_construct(0.2, 0.4, 0.6, 0.8) as doc_vec ) """) # 2. 初始化并拟合TruncatedSVD from snowflake.ml.modeling.decomposition import TruncatedSVD tsvd = TruncatedSVD( input_cols=[f"vec_{i}" for i in range(vec_length)], output_cols="out_svd", n_components=2 ) tsvd.fit(df_expanded) result = tsvd.transform(df_expanded) result.show()
解决方案二:稀疏数组转稀疏矩阵(适合大尺寸稀疏嵌入)
针对尺寸>1000的稀疏数组,展开多列会导致性能问题,可通过自定义模型包装sklearn的TruncatedSVD,直接处理稀疏矩阵:
from snowflake.snowpark.functions import udf from scipy.sparse import csr_matrix from snowflake.ml.modeling.framework import BaseTransformer from sklearn.decomposition import TruncatedSVD as SklearnTSVD # 自定义Transformer处理稀疏数组 class SparseTSVD(BaseTransformer): def __init__(self, input_col: str, output_col: str, n_components: int=2): self.input_col = input_col self.output_col = output_col self.tsvd = SklearnTSVD(n_components=n_components) self.fitted = False def fit(self, dataset): # 将Snowpark数组转为scipy稀疏矩阵 arr_rows = dataset.select(self.input_col).collect() sparse_matrix = csr_matrix([row[self.input_col] for row in arr_rows]) self.tsvd.fit(sparse_matrix) self.fitted = True return self def transform(self, dataset): if not self.fitted: raise ValueError("模型未完成拟合,请先调用fit方法") arr_rows = dataset.select(self.input_col).collect() sparse_matrix = csr_matrix([row[self.input_col] for row in arr_rows]) transformed_arr = self.tsvd.transform(sparse_matrix) # 将降维结果转为Snowpark数组列 return dataset.with_column(self.output_col, [list(row) for row in transformed_arr]) # 使用自定义模型处理数据 tsvd_sparse = SparseTSVD(input_col="doc_vec", output_col="out_svd", n_components=2) tsvd_sparse.fit(df) result_sparse = tsvd_sparse.transform(df) result_sparse.show()
注意事项
- 大稀疏数组优先选择方案二,避免展开多列带来的存储和计算开销
- Snowpark ML原生
TruncatedSVD暂不支持直接输入ARRAY列,必须转换为多数值列或稀疏矩阵格式 - 若数据量极大,自定义模型中的
collect()会将数据拉到本地,可考虑结合Snowflake的分布式UDF优化处理逻辑
内容的提问来源于stack exchange,提问作者user2148414
相关产品推荐
相关产品推荐

