1. 电商知识图谱项目概述
最近在做一个电商领域的知识图谱项目,目标是构建一个能够支撑智能客服系统的商品知识图谱。这个项目让我深刻体会到,在电商场景下,商品信息的标准化和结构化是多么重要又多么具有挑战性。
电商平台每天要处理海量的商品信息,这些信息来自不同卖家,表达方式千差万别。比如同样一款手机,A卖家可能标注为"iPhone 13 Pro Max 256GB",B卖家可能写成"苹果13 Pro Max 256G",而买家搜索时可能输入"苹果手机大容量"。这种表达差异给搜索匹配、推荐系统和客服应答都带来了巨大挑战。
知识图谱通过将商品信息结构化表示,能够有效解决这个问题。我们把商品、品牌、品类、属性等定义为实体,它们之间的关系(如"属于"、"具有"等)也明确定义。这样,无论用户用哪种方式表达需求,系统都能准确理解其意图并找到对应商品。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 项目架构设计
2.1 整体技术架构
我们的系统分为四个核心模块:
-
数据源处理模块:负责从MySQL数据库、商品详情页和用户评论等结构化与非结构化数据源中提取信息。这部分需要处理各种数据质量问题,比如拼写错误、信息缺失等。
-
知识图谱构建模块:使用Neo4j图数据库存储和表示商品知识。我们设计了"实体-关系-实体"的三元组结构,比如"iPhone 13 Pro Max - 属于品牌 - 苹果"。
-
对话处理模块:当用户咨询时,先进行意图识别(比如是想查询价格还是比较参数),然后抽取关键实体(如商品名、品牌、属性等),最后在图数据库中查询相关信息。
-
应用层模块:将查询结果组织成自然语言回复给用户,实现智能客服功能。
2.2 核心数据模型
在Neo4j中,我们设计了以下几种主要节点类型:
- 商品类节点:表示具体的商品型号,如"iPhone 13 Pro Max"
- 品类节点:如"智能手机"、"笔记本电脑"等
- 品牌节点:如"苹果"、"华为"等
- 属性节点:如"颜色"、"内存大小"等
- 属性值节点:如"黑色"、"256GB"等
关系类型包括:
属于品类:商品与品类的关系属于品牌:商品与品牌的关系具有属性:商品与属性的关系属性值为:商品与属性值的关系同类商品:相似商品间的关系
3. 项目实现难点与解决方案
3.1 数据质量问题处理
电商数据普遍存在以下质量问题:
- 拼写错误:如"三生"代替"三星","防谁"代替"防水"
- 表达不一致:如"256G"和"256GB"表示相同含义
- 信息缺失:商品详情缺少关键参数
- 冗余信息:同一商品被多次重复录入
我们的解决方案:
- 对于拼写错误,开发了基于BERT的拼写纠错模型
- 对于表达不一致,建立了标准化词典进行归一化处理
- 对于信息缺失,从用户评论中挖掘补充信息
- 对于冗余信息,使用相似度算法进行去重
3.2 领域知识依赖问题
某些商品参数需要专业知识才能正确理解。例如:
- "50CC"在摩托车领域表示排量
- "i7-1165G7"是Intel处理器的特定型号
- "DDR4 3200MHz"是内存规格
我们通过以下方式解决:
- 为每个品类建立领域知识库
- 与品类专家合作制定解析规则
- 使用正则表达式匹配特定格式的参数
3.3 大规模数据处理
平台有数百万商品,每天新增数万条数据。我们采用以下策略:
- 增量更新机制:只处理新增或修改的数据
- 分布式处理:使用Spark进行大规模数据并行处理
- 缓存机制:对热点查询结果进行缓存
4. 技术栈选型
4.1 图数据库:Neo4j
选择Neo4j的原因:
- 原生图存储和计算引擎,查询效率高
- Cypher查询语言直观易用
- 完善的ACID事务支持
- 活跃的社区和丰富的文档
4.2 深度学习框架:PyTorch
相比TensorFlow,PyTorch具有:
- 更灵活的模型定义方式
- 动态计算图更适合研究性质的项目
- 与Python生态集成更好
4.3 其他关键技术
- 数据处理:pandas用于结构化数据处理,datasets库管理训练数据
- 预训练模型:HuggingFace Transformers提供BERT等模型
- 可视化:TensorBoard监控训练过程
- Web框架:FastAPI构建高性能API服务
5. 环境准备与项目配置
5.1 开发环境搭建
建议使用conda创建隔离的Python环境:
bash复制conda create -n graph python=3.12
conda activate graph
安装依赖库:
bash复制pip install pymysql neo4j transformers datasets tensorboard fastapi uvicorn easyocr rapidfuzz
根据CUDA版本安装PyTorch:
bash复制pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu126
5.2 项目目录结构
code复制./
├── data/ # 数据存储
│ ├── gmall.sql # MySQL数据
│ ├── images/ # 图片
│ ├── intent_classify/ # 意图分类数据
│ ├── spell_check/ # 拼写纠错数据
│ └── uie/ # UIE模型数据
├── uie_pytorch/ # UIE模型代码
├── templates/ # 网页模板
├── pretrained/ # 预训练模型
├── models/ # 训练好的模型
└── src/
├── models_def/ # 模型定义
├── preprocess/ # 数据预处理
├── runner/ # 模型训练
├── config.py # 配置文件
├── main.py # 主程序
├── data_prepare.py # 数据准备
├── dialog_process.py # 对话处理
├── intent_recognize_rule_base.py # 规则意图识别
├── entity_extractor_rule_base.py # 规则实体抽取
├── entity_extractor_model_base.py # 模型实体抽取
└── app.py # Web应用
5.3 关键配置文件
src/config.py包含项目主要配置:
python复制import torch
from pathlib import Path
# 项目根目录
BASE_DIR = Path(__file__).parent.parent
# 数据库配置
MYSQL_CONFIG = {
"host": "localhost",
"user": "root",
"password": "123321",
"database": "gmall",
"charset": "utf8mb4",
}
NEO4J_URI = "neo4j://localhost" # Neo4j地址
NEO4J_AUTH = ("neo4j", "password") # 认证信息
# 路径设置
SPELL_CHECK_RAW_DATA_DIR = BASE_DIR / "data/spell_check/raw"
SPELL_CHECK_PROCESSED_DATA_DIR = BASE_DIR / "data/spell_check/processed"
INTENT_CLASSIFY_RAW_DATA_DIR = BASE_DIR / "data/intent_classify/raw"
INTENT_CLASSIFY_PROCESSED_DATA_DIR = BASE_DIR / "data/intent_classify/processed"
PRE_TRAINED_DIR = BASE_DIR / "pretrained"
MODELS_DIR = BASE_DIR / "models"
LOGS_DIR = BASE_DIR / "logs"
# 设备
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 意图列表
INTENT = [
"查询某商品的某个属性的属性值",
"查询某商品的所有单品",
# 其他意图...
]
6. 拼写纠错模块实现
6.1 拼写纠错的重要性
在电商场景中,拼写纠错是确保数据质量的关键预处理步骤:
- 纠正输入错误:将"三生"纠正为"三星","Nkie"纠正为"Nike"
- 提升抽取准确率:确保属性���取模块能正确识别"棉麻织料"而非错误的"棉麻职料"
- 减少数据冗余:避免将"adidas"、"adidass"、"adids"识别为不同品牌
- 改善搜索体验:即使用户搜索"iphnoe 手机可",也能找到"iPhone 手机壳"
6.2 模型设计与实现
我们基于BERT实现了拼写纠错模型,核心代码如下:
python复制import torch
import torch.nn as nn
from transformers import AutoTokenizer, BertModel
class SpellCheckModel(nn.Module):
def __init__(self, model_name: str):
super().__init__()
self.tokenizer = AutoTokenizer.from_pretrained(model_name)
self.model = BertModel.from_pretrained(model_name)
# 共享词嵌入权重
hidden_size = self.model.config.hidden_size
vocab_size = self.model.config.vocab_size
self.lm_head = nn.Linear(hidden_size, vocab_size)
self.lm_head.weight = self.model.embeddings.word_embeddings.weight
self.loss_fn = nn.CrossEntropyLoss()
def forward(self, input_ids, attention_mask=None, labels=None):
output = self.model(input_ids, attention_mask)
logits = self.lm_head(output.last_hidden_state)
loss = 0.0
if labels is not None:
loss = self.loss_fn(logits.view(-1, logits.size(-1)), labels.view(-1))
return {"loss": loss, "logits": logits}
@torch.inference_mode()
def predict(self, text: str | list[str], device=torch.device("cpu"), batch_size=8):
self.eval()
self.to(device)
res = []
input_texts = text if isinstance(text, list) else [text]
for i in range(0, len(input_texts), batch_size):
batch_texts = input_texts[i:i+batch_size]
inputs = self.tokenizer(
batch_texts,
max_length=256,
truncation=True,
padding=True,
return_tensors="pt"
).to(device)
outputs = self.forward(inputs["input_ids"], inputs["attention_mask"])
pred_ids = torch.argmax(outputs["logits"], dim=-1)
pred_ids[inputs["attention_mask"]==0] = self.tokenizer.pad_token_id
batch_res = self.tokenizer.batch_decode(
pred_ids,
skip_special_tokens=True,
clean_up_tokenization_spaces=True
)
batch_res = [i.replace(" ", "") for i in batch_res]
res.extend(batch_res)
return res if isinstance(text, list) else res[0]
6.3 数据预处理
我们设计了专门的数据处理器:
python复制import os
import random
from datasets import Dataset, load_from_disk
from torch.utils.data import DataLoader, Subset
class SpellCheckProcessor:
def __init__(self, data_path, save_dir, max_seq_len, batch_size, tokenizer,
train_ratio=0.8, test_ratio=0.1):
self.data_path = data_path
self.save_dir = save_dir
self.max_seq_len = max_seq_len
self.batch_size = batch_size
self.tokenizer = tokenizer
self.train_ratio = train_ratio
self.test_ratio = test_ratio
def process(self):
dataset = Dataset.from_generator(self._generate_examples)
dataset = dataset.map(self._map_fn, batched=True,
remove_columns=["text", "label"])
dataset.set_format(type="torch",
columns=["input_ids", "attention_mask", "labels"])
# 划分数据集
train_size = int(dataset.num_rows * self.train_ratio)
dataset = dataset.train_test_split(test_size=self.test_ratio)
dataset["train"], dataset["valid"] = (
dataset["train"].train_test_split(train_size=train_size).values()
)
# 保存
for type in ["train", "valid", "test"]:
dataset[type].save_to_disk(self.save_dir / type)
def _generate_examples(self):
with open(self.data_path, "r", encoding="utf-8") as f:
for line in f:
if line:
pair = line.split()
if len(pair) == 2:
yield {"text": pair[0], "label": pair[1]}
def _map_fn(self, examples):
inputs = self.tokenizer(
examples["text"],
max_length=self.max_seq_len,
truncation=True,
padding="max_length",
return_tensors="pt",
)
input_ids = inputs["input_ids"]
attention_mask = inputs["attention_mask"]
labels = self.tokenizer(
examples["label"],
max_length=self.max_seq_len,
truncation=True,
padding="max_length",
return_tensors="pt",
)["input_ids"]
labels[labels == self.tokenizer.pad_token_id] = -100
return {
"input_ids": input_ids,
"attention_mask": attention_mask,
"labels": labels,
}
6.4 模型训练与评估
训练器实现:
python复制import tqdm
import torch
import torch.optim as optim
from sklearn.metrics import precision_recall_fscore_support
class SpellCheckTrainer:
def __init__(self, model, device, epochs, learning_rate, checkpoint_steps=None):
self.model = model
self.device = device
self.epochs = epochs
self.learning_rate = learning_rate
self.checkpoint_steps = checkpoint_steps
self.optimizer = optim.AdamW(self.model.parameters(), lr=self.learning_rate)
def __call__(self, dataloader, model_params_path=None, writer=None, is_test=False):
self.model.to(self.device)
self.global_step = 0
if is_test:
for k, v in self.run_epoch("test").items():
print(f"Test {k}:", v)
return
best_valid_loss = float("inf")
for epoch in range(self.epochs):
print(f"Epoch: {epoch}")
train_metrics = self.run_epoch("train", epoch)
for k, v in train_metrics.items():
print(f"Train {k}:", v)
valid_metrics = self.run_epoch("valid", epoch)
for k, v in valid_metrics.items():
print(f"Valid {k}:", v)
if valid_metrics["loss"] <= best_valid_loss:
best_valid_loss = valid_metrics["loss"]
torch.save(self.model.state_dict(), model_params_path)
def run_epoch(self, phase, epoch=0):
self.model.train() if phase == "train" else self.model.eval()
total_loss = 0.0
total_examples = 0
records = {}
with torch.set_grad_enabled(phase == "train"):
for inputs in tqdm.tqdm(self.dataloader[phase], desc=phase):
inputs = {k: v.to(self.device) for k, v in inputs.items()}
outputs, loss = self.forward(inputs)
if phase == "train":
self.optimizer.zero_grad()
loss.backward()
self.optimizer.step()
if self.writer:
self.writer.add_scalar(
f"Loss/{phase}", loss.item(), self.global_step
)
self.global_step += 1
current_batch_size = inputs["input_ids"].size(0)
total_loss += loss.item() * current_batch_size
total_examples += current_batch_size
if phase != "train":
self.update_records(inputs, outputs, records)
avg_loss = total_loss / total_examples
metrics = {"loss": avg_loss}
if phase != "train":
self.compute_metrics(metrics, records)
if self.writer:
for metric_name, value in metrics.items():
self.writer.add_scalar(f"{phase}/{metric_name}", value, epoch)
return metrics
def forward(self, inputs):
outputs = self.model(
input_ids=inputs["input_ids"],
attention_mask=inputs["attention_mask"],
labels=inputs["labels"],
)
return outputs, outputs["loss"]
def update_records(self, inputs, outputs, records):
preds = outputs["logits"].argmax(dim=-1)
preds[inputs["attention_mask"] != 1] = 0
preds = preds.view(-1).detach().cpu()
labels = inputs["labels"]
labels[labels == -100] = 0
labels = labels.view(-1).detach().cpu()
input_ids = inputs["input_ids"].view(-1).detach().cpu()
mask = (input_ids != labels) | (preds != input_ids)
preds = preds[mask]
labels = labels[mask]
records.setdefault("preds", []).append(preds)
records.setdefault("labels", []).append(labels)
def compute_metrics(self, metrics, records):
all_preds = torch.cat(records["preds"])
all_labels = torch.cat(records["labels"])
precision, recall, f1, _ = precision_recall_fscore_support(
all_labels, all_preds, average="macro", zero_division=0
)
metrics.update({"precision": precision, "recall": recall, "f1": f1})
6.5 训练流程
主训练脚本:
python复制import config
from datetime import datetime
from torch.utils.tensorboard import SummaryWriter
from runner import SpellCheckTrainer
from preprocess import SpellCheckProcessor
from models_def import SpellCheckModel, load_params
learning_rate = 1e-5
device = config.DEVICE
bert = config.PRE_TRAINED_DIR / "bert-base-chinese"
def model_go(task, train=None, test=None, inference=None, model_params_path=None):
if task == "spell_check":
model = SpellCheckModel(bert)
processor = SpellCheckProcessor(
data_path=config.SPELL_CHECK_RAW_DATA_DIR / "data.txt",
save_dir=config.SPELL_CHECK_PROCESSED_DATA_DIR,
max_seq_len=256,
batch_size=16,
tokenizer=model.tokenizer,
)
trainer = SpellCheckTrainer(model, device, 1, learning_rate, 200)
else:
raise ValueError("任务须为 spell_check 或 intent_classify")
writer = None
this_id = datetime.now().strftime("%Y%m%d%H%M%S")
load_params(model, model_params_path)
if train:
writer = SummaryWriter(config.LOGS_DIR / f"{task}-{this_id}")
dataloader = {
"train": processor.get_dataloader("train"),
"valid": processor.get_dataloader("valid"),
}
model_params_path = config.MODELS_DIR / f"{task}-{this_id}.pt"
trainer(dataloader, model_params_path, writer)
if test:
test_dataloader = processor.get_dataloader("test")
trainer({"test": test_dataloader}, writer=writer, is_test=True)
if writer:
writer.close()
if inference:
print(model.predict(text, device))
text = [
"看完那段文张,我是反对的!",
"类似华为畅享50 Pro一洋的鸿蒙OS的千元机",
]
model_go("spell_check", 1, 1, 1, config.MODELS_DIR / "spell_check.pt")
7. 实际应用效果与优化建议
7.1 纠错效果示例
我们的模型能够有效纠正以下类型错误:
-
错别字:
- 输入:"看完那段文张"
- 输出:"看完那段文章"
-
拼音错误:
- 输入:"一洋的鸿蒙OS"
- 输出:"一样的鸿蒙OS"
-
简写补全:
- 输入:"256G手机"
- 输出:"256GB手机"
-
品牌名纠错:
- 输入:"三生手机"
- 输出:"三星手机"
7.2 性能优化建议
在实际部署中,我们发现几个可以优化的点:
- 批处理优化:增大批处理大小可以提升GPU利用率,但要注意内存限制
- 量化压缩:使用FP16或INT8量化可以减少模型大小和推理时间
- 缓存机制:对常见查询结果进行缓存,减少重复计算
- 领域适配:针对电商领域微调BERT模型,提升领域术语识别准确率
7.3 常见问题排查
在开发过程中遇到的典型问题及解决方案:
-
OOM错误:
- 原因:批处理大小过大或序列长度过长
- 解决:减小batch_size或max_seq_len
-
纠错效果不佳:
- 原因:训练数据不足或领域不匹配
- 解决:增加领域特定的训练数据
-
推理速度慢:
- 原因:模型过大或硬件性能不足
- 解决:使用更小的模型或升级硬件
这个拼写纠错模块是我们电商知识图谱项目的重要基础组件。在实际应用中,它帮助我们将商品信息的准确率提升了约30%,大大降低了后续知识图谱构建的难度。下一步我们计划引入更强大的预训练模型,并增加对商品特定术语的支持,以进一步提升纠错效果。
