bert_base/tools/product_name/Raner/product_info.py

327 lines
10 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import pandas as pd
import sys
from modelscope.pipelines import pipeline
from modelscope.utils.constant import Tasks
import re
# 定义需要跳过的标签(这些标签的商品不进行推理)
SKIP_LABELS = ['当当图书']
# 全局 NER pipeline
ner_pipeline = None
# 品牌缓存 {品牌关键词: 完整品牌名}
brand_cache = {}
# 缓存命中统计
cache_hit_count = 0
cache_miss_count = 0
def load_model():
"""
加载 ModelScope NER 模型
"""
global ner_pipeline
print("🔧 加载 ModelScope NER 模型中...")
ner_pipeline = pipeline(
Tasks.named_entity_recognition,
'damo/nlp_raner_named-entity-recognition_chinese-base-ecom-50cls'
)
print("✅ 模型加载完成!")
return ner_pipeline
def merge_brand_names(brand_spans):
"""
合并多个品牌名称
规则:如果有中文和英文,拼接为 中文(英文)
如果只有一个,直接返回
"""
if not brand_spans:
return None
# 分离中文和英文
chinese_brands = []
english_brands = []
other_brands = []
for brand in brand_spans:
# 判断是否包含中文字符
if re.search(r'[\u4e00-\u9fff]', brand):
chinese_brands.append(brand)
# 判断是否为英文(只包含英文字母、数字、空格、.、-
elif re.match(r'^[a-zA-Z0-9\s\.\-]+$', brand):
english_brands.append(brand)
else:
other_brands.append(brand)
# 构建最终品牌名
result_parts = []
# 如果有中文品牌
if chinese_brands:
chinese_name = chinese_brands[0]
# 如果有多个中文品牌,用 / 连接
if len(chinese_brands) > 1:
chinese_name = '/'.join(chinese_brands)
result_parts.append(chinese_name)
# 如果有英文品牌
if english_brands:
english_name = english_brands[0]
if len(english_brands) > 1:
english_name = '/'.join(english_brands)
# 如果已经有中文,用括号包裹英文
if result_parts:
result_parts.append(f"({english_name})")
else:
result_parts.append(english_name)
# 如果有其他品牌(非中英文)
if other_brands:
for brand in other_brands:
if brand not in result_parts:
result_parts.append(brand)
return ''.join(result_parts)
def extract_brand_spans(ner_result):
"""
从 ModelScope NER 结果中提取所有品牌 span
返回: [品牌span列表]
"""
brands = []
if isinstance(ner_result, dict):
output_list = ner_result.get('output', [])
for item in output_list:
if item.get('type') == '品牌':
span = item.get('span', '')
if span:
brands.append(span)
return brands
def check_cache(text):
"""
检查缓存中是否有匹配的品牌
遍历缓存如果缓存中的key包含在商品名称中则返回对应的品牌
返回: (是否命中, 品牌名)
"""
global cache_hit_count, cache_miss_count
# 按key长度从长到短排序优先匹配更长的关键词更精确
sorted_keys = sorted(brand_cache.keys(), key=len, reverse=True)
for key in sorted_keys:
if key in text:
cache_hit_count += 1
print(f" ✅ 缓存命中: '{key}' -> '{brand_cache[key]}'")
return True, brand_cache[key]
cache_miss_count += 1
return False, None
def update_cache(product_text, brand_spans, merged_brand):
"""
更新缓存
直接将模型返回的品牌 span 作为 keymerged_brand 作为 value
"""
global brand_cache
if not merged_brand or merged_brand == "其他":
return
for span in brand_spans:
if span and span not in brand_cache:
brand_cache[span] = merged_brand
print(f" 💾 缓存更新: '{span}' -> '{merged_brand}'")
def reco_single(text, use_cache=True):
"""
对单个商品名称进行推理,提取品牌信息
支持缓存机制
"""
global brand_cache
if ner_pipeline is None:
raise RuntimeError("模型未加载,请先调用 load_model()")
# 1. 先检查缓存
if use_cache and brand_cache:
hit, cached_brand = check_cache(text)
if hit:
return {'品牌': [[cached_brand, 1.0]]}
# 2. 缓存未命中,调用模型
try:
ner_result = ner_pipeline(text)
brand_spans = extract_brand_spans(ner_result)
if not brand_spans:
return {'品牌': []}
# 3. 合并品牌名称
merged_brand = merge_brand_names(brand_spans)
if merged_brand:
# 4. 更新缓存:用模型返回的 span 作为 key
update_cache(text, brand_spans, merged_brand)
return {'品牌': [[merged_brand, 1.0]]}
else:
return {'品牌': []}
except Exception as e:
print(f"⚠️ 推理失败: {text} - {e}")
return {'品牌': []}
def extract_brand_from_ner(ner_tags):
"""
从NER推理结果中提取品牌
返回: {商品名称: 品牌名称}
"""
result = {}
for product_name, tags in ner_tags.items():
brand_list = tags.get('品牌', [])
if not brand_list:
brand_name = "其他"
else:
# 取第一个品牌(已经合并过的)
brand_name = brand_list[0][0]
result[product_name] = brand_name
return result
def should_skip_inference(label_value):
"""
判断是否应该跳过推理
"""
if pd.isna(label_value):
return False
return str(label_value).strip() in SKIP_LABELS
def batch_process(input_file='input.xlsx', output_file='output.xlsx',
key_column='sentence', label_column='label', batch_size=32):
"""
批量处理大文件
"""
global cache_hit_count, cache_miss_count, brand_cache
try:
print(f"📂 读取文件: {input_file}")
df = pd.read_excel(input_file, engine='openpyxl')
if key_column not in df.columns:
print(f"❌ 错误: 找不到列 '{key_column}'")
sys.exit(1)
if label_column not in df.columns:
print(f"⚠️ 警告: 找不到列 '{label_column}',将创建新列")
df[label_column] = ""
# 数据清洗
df = df.dropna(subset=[key_column])
df[key_column] = df[key_column].astype(str).str.strip()
df[label_column] = df[label_column].fillna("").astype(str).str.strip()
# 分离需要推理和不需要推理的数据
print("🔍 检查标签,筛选需要推理的数据...")
skip_mask = df[label_column].apply(should_skip_inference)
need_inference_mask = ~skip_mask
df_need_inference = df[need_inference_mask].copy()
df_skip = df[skip_mask].copy()
print(f" 📊 总共 {len(df)} 行数据")
print(f" ✅ 需要推理: {len(df_need_inference)}")
print(f" ⏭️ 跳过推理: {len(df_skip)}")
# 处理需要推理的数据
if len(df_need_inference) > 0:
inference_data = df_need_inference[key_column].tolist()
print(f"🤖 开始逐条推理 (批次大小: {batch_size})...")
print(f"📝 当前缓存条目数: {len(brand_cache)}")
all_ner_tags = {}
total = len(inference_data)
for i, text in enumerate(inference_data):
if (i + 1) % batch_size == 0 or i == total - 1:
print(
f" 进度: {i + 1}/{total} | 缓存命中: {cache_hit_count} | 未命中: {cache_miss_count} | 缓存条目: {len(brand_cache)}")
ner_result = reco_single(text, use_cache=True)
all_ner_tags[text] = ner_result
# 提取品牌
print("📊 提取品牌信息...")
brand_dict = extract_brand_from_ner(all_ner_tags)
# 添加品牌列
df_need_inference['brand'] = df_need_inference[key_column].map(brand_dict)
df_need_inference['brand'] = df_need_inference['brand'].fillna("其他")
# 处理跳过的数据
if len(df_skip) > 0:
df_skip['brand'] = df_skip[label_column]
# 合并数据并恢复原始顺序
df_result = pd.concat([df_need_inference, df_skip], ignore_index=True)
df_result = df_result.sort_index().reset_index(drop=True)
# 保存结果
df_result.to_excel(output_file, index=False, engine='openpyxl')
print(f"✅ 处理完成!结果已保存到: {output_file}")
# 统计信息
total = len(df_result)
found = (df_result['brand'] != "其他").sum()
skipped_count = len(df_result[df_result[label_column].apply(should_skip_inference)])
print(f"\n📊 统计信息:")
print(f" 总行数: {total}")
print(f" 识别品牌: {found}")
print(f" 跳过推理: {skipped_count}")
print(f" 缓存命中: {cache_hit_count}")
print(f" 缓存未命中: {cache_miss_count}")
if (cache_hit_count + cache_miss_count) > 0:
print(f" 缓存命中率: {cache_hit_count / (cache_hit_count + cache_miss_count) * 100:.1f}%")
print(f" 缓存条目数: {len(brand_cache)}")
print(f"\n📋 结果预览 (前10行):")
print(df_result[[key_column, label_column, 'brand']].head(10))
# 打印缓存内容供参考
if brand_cache:
print(f"\n📝 缓存内容:")
for k, v in brand_cache.items():
print(f" '{k}' -> '{v}'")
return df_result
except FileNotFoundError:
print(f"❌ 错误: 找不到文件 '{input_file}'")
sys.exit(1)
except Exception as e:
print(f"❌ 处理失败: {e}")
import traceback
traceback.print_exc()
sys.exit(1)
if __name__ == "__main__":
# 加载模型
load_model()
# 执行批量处理
batch_process('input.xlsx', 'output.xlsx', 'sentence', 'label', batch_size=32)