使用tidymodels+LIME做文本预测时的绘图异常与报错问题
使用LIME解释tidymodels拟合的LASSO文本模型的问题解决
问题说明
用lime包解释tidymodels构建的LASSO文本分类模型时,碰到两个问题:
plot_features(explanation)生成的图形不是LIME标准的特征贡献图plot_text_explanations(explanation)报错:Error: original_text is not a string (a length one character vector),无法生成文本高亮解释图
要求仅使用tidymodels工具,不涉及caret等其他包,原始可复现代码如下:
library("conflicted") library("lime") conflict_prefer("explain", "lime") library("textrecipes") library("tidymodels") library("tidyverse") conflict_prefer("slice", "dplyr") reviews <- quanteda.textmodels::data_corpus_moviereviews |> quanteda::convert(to = "data.frame") |> tibble() |> select(rating = sentiment, text) set.seed(1234) review_split <- initial_split(reviews, strata = rating) review_train <- training(review_split) review_test <- testing(review_split) lasso_recipe <- recipe(rating ~ text, data = review_train) |> step_tokenize(text) |> step_stopwords(text) |> step_tokenfilter(text, max_tokens = 100) |> step_tfidf(text) |> step_normalize(all_predictors()) lasso_spec <- logistic_reg(penalty = 0.1, mixture = 1) |> set_mode("classification") |> set_engine("glmnet") lasso_wf <- workflow() |> add_recipe(lasso_recipe) |> add_model(lasso_spec) lasso_fit <- lasso_wf |> fit(data = review_train) predict(lasso_fit, review_test) #> # A tibble: 500 × 1 #> .pred_class #> <fct> #> 1 neg #> 2 neg #> 3 neg #> 4 pos #> 5 neg #> 6 neg #> 7 pos #> 8 pos #> 9 pos #> 10 neg #> # … with 490 more rows preprocess <- function(input) { baked <- recipe(rating ~ text, data = input) |> step_tokenize(text) |> step_stopwords(text) |> step_tokenfilter(text, max_tokens = 100) |> step_tfidf(text) |> step_normalize(all_predictors()) |> prep() |> bake(new_data = NULL) |> select(-rating) return(baked) } preprocess(slice(reviews, 1:3)) #> # A tibble: 3 × 100 #> tfidf_text_10 tfidf_…¹ tfidf…² tfidf…³ tfidf…⁴ tfidf…⁵ tfidf…⁶ tfidf…⁷ tfidf…⁸ #> <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> #> 1 1.15 -0.577 1.15 1.15 -0.813 1.15 0.283 1.15 -0.577 #> 2 -0.577 -0.577 -0.577 -0.577 1.12 -0.577 -1.11 -0.577 1.15 #> 3 -0.577 1.15 -0.577 -0.577 -0.304 -0.577 0.828 -0.577 -0.577 #> # … with 91 more variables: tfidf_text_attempt <dbl>, #> # tfidf_text_audience <dbl>, tfidf_text_away <dbl>, tfidf_text_back <dbl>, #> # tfidf_text_bad <dbl>, tfidf_text_baldwin <dbl>, tfidf_text_based <dbl>, #> # tfidf_text_big <dbl>, tfidf_text_biggest <dbl>, #> # tfidf_text_characters <dbl>, tfidf_text_chase <dbl>, #> # tfidf_text_claire <dbl>, tfidf_text_clear <dbl>, tfidf_text_comes <dbl>, #> # tfidf_text_coming <dbl>, tfidf_text_cool <dbl>, tfidf_text_course <dbl>, … explainer <- lime( slice(reviews, 1:3), model = extract_fit_parsnip(lasso_fit), preprocess = preprocess ) explanation <- explain( slice(reviews, 1:3), explainer = explainer, labels = "pos", n_features = 10 ) # plot_features() does not produce the correct plot plot_features(explanation)

# plot_text_explanations() issues an error plot_text_explanations(explanation) #> Error: original_text is not a string (a length one character vector).
问题解决
1. 修复plot_features图形异常
问题根源:自定义preprocess函数重新创建配方,导致特征映射和训练模型不一致;同时lime初始化时未明确分类类型,且输入数据格式不对。
修正步骤:
- 复用训练好的配方进行预处理,避免特征差异
- 初始化
explainer时传入原始文本向量,指定type = "classification" - 调用
explain时开启case_level = TRUE,确保每个样本的特征贡献单独展示
修改后的预处理函数和解释器代码:
# 复用训练好的lasso_recipe,避免重新构建导致特征不匹配 preprocess <- function(input) { bake(prep(lasso_recipe), new_data = input) %>% select(-rating) } # 初始化解释器时直接传入原始文本向量 explainer <- lime( x = review_train$text, model = extract_fit_parsnip(lasso_fit), preprocess = preprocess, type = "classification" ) # 解释时传入文本向量,而非数据框 explanation <- explain( x = slice(reviews, 1:3)$text, explainer = explainer, labels = "pos", n_features = 10, case_level = TRUE )
此时运行plot_features(explanation)就能生成标准的LIME特征贡献图,每个样本的正向/负向特征会清晰展示。
2. 修复plot_text_explanations报错
问题根源:plot_text_explanations需要解释对象包含单字符格式的原始文本,之前传入数据框导致LIME无法正确提取文本内容。
修正后的代码已在上面给出,确保explain传入的是文本向量而非数据框。此时运行:
plot_text_explanations(explanation)
就能生成文本高亮解释图,影响分类的关键词会用不同颜色标注(正向贡献为绿色,负向贡献为红色)。
完整可运行修正代码
library("conflicted") library("lime") conflict_prefer("explain", "lime") library("textrecipes") library("tidymodels") library("tidyverse") conflict_prefer("slice", "dplyr") reviews <- quanteda.textmodels::data_corpus_moviereviews |> quanteda::convert(to = "data.frame") |> tibble() |> select(rating = sentiment, text) set.seed(1234) review_split <- initial_split(reviews, strata = rating) review_train <- training(review_split) review_test <- testing(review_split) lasso_recipe <- recipe(rating ~ text, data = review_train) |> step_tokenize(text) |> step_stopwords(text) |> step_tokenfilter(text, max_tokens = 100) |> step_tfidf(text) |> step_normalize(all_predictors()) lasso_spec <- logistic_reg(penalty = 0.1, mixture = 1) |> set_mode("classification") |> set_engine("glmnet") lasso_wf <- workflow() |> add_recipe(lasso_recipe) |> add_model(lasso_spec) lasso_fit <- lasso_wf |> fit(data = review_train) # 修正预处理函数:复用训练好的配方 preprocess <- function(input) { bake(prep(lasso_recipe), new_data = input) %>% select(-rating) } # 修正解释器初始化:传入原始文本向量,指定分类类型 explainer <- lime( x = review_train$text, model = extract_fit_parsnip(lasso_fit), preprocess = preprocess, type = "classification" ) # 修正解释调用:传入文本向量,开启case_level explanation <- explain( x = slice(reviews, 1:3)$text, explainer = explainer, labels = "pos", n_features = 10, case_level = TRUE ) # 生成正确的特征贡献图 plot_features(explanation) # 生成文本高亮解释图 plot_text_explanations(explanation)
内容的提问来源于stack exchange,提问作者captain
相关产品推荐
相关产品推荐

