如何训练本地模型实现发票文本转JSON及BERT调优排查
一、本地模型实现发票文本转指定JSON(无OpenAI依赖)
1. 模型选型
优先选择Seq2Seq类预训练模型(如T5、BART),这类模型天生适配“文本输入→结构化文本输出”的场景,无需依赖外部API,可完全本地训练部署。推荐从轻量版模型入手(如t5-small、facebook/bart-base),平衡性能与硬件需求。
2. 数据准备
将发票文本作为模型输入,目标结构化JSON转为字符串作为输出,需保证格式统一:
示例数据
- 输入文本:
Amazon.com Invoice
Order Date: 05/15/2024
Order #123-4567890-1234567
Item: Wireless Headphones - $89.99
Shipping: $5.99
Tax: $7.20
Total: $103.18
- 目标输出JSON字符串:
{"order_number": "123-4567890-1234567", "order_date": "05/15/2024", "total_amount": 103.18, "items": [{"name": "Wireless Headphones", "price": 89.99}], "shipping": 5.99, "tax": 7.20}
3. 训练核心代码
使用Hugging Face Transformers & Trainer API实现本地训练:
from transformers import T5Tokenizer, T5ForConditionalGeneration, Seq2SeqTrainingArguments, Seq2SeqTrainer import datasets # 加载数据集(假设已整理为csv格式,包含input_text和target_json列) dataset = datasets.load_dataset('csv', data_files='amazon_invoices.csv') # 加载模型与Tokenizer tokenizer = T5Tokenizer.from_pretrained('t5-small') model = T5ForConditionalGeneration.from_pretrained('t5-small') # 数据预处理函数 def preprocess_function(examples): inputs = tokenizer(examples['input_text'], max_length=512, truncation=True, padding='max_length') targets = tokenizer(examples['target_json'], max_length=256, truncation=True, padding='max_length') inputs['labels'] = targets['input_ids'] return inputs tokenized_dataset = dataset.map(preprocess_function, batched=True) # 设置训练参数 training_args = Seq2SeqTrainingArguments( output_dir='./invoice_model', per_device_train_batch_size=8, per_device_eval_batch_size=8, num_train_epochs=5, learning_rate=3e-5, weight_decay=0.01, logging_dir='./logs', logging_steps=10, evaluation_strategy='epoch', save_strategy='epoch' ) # 初始化Trainer trainer = Seq2SeqTrainer( model=model, args=training_args, train_dataset=tokenized_dataset['train'], eval_dataset=tokenized_dataset['test'] ) # 开始训练 trainer.train() # 推理示例 def generate_invoice_json(text): inputs = tokenizer(text, return_tensors='pt', truncation=True, max_length=512).to(model.device) outputs = model.generate(**inputs, max_length=256) json_str = tokenizer.decode(outputs[0], skip_special_tokens=True) # 可选:添加JSON格式校验与修复 import json try: return json.loads(json_str) except: return {"error": "Invalid JSON generated", "raw_output": json_str}
4. 推理优化
训练后可添加JSON格式校验(如json.loads捕获异常),对生成的字符串做简单修复(如补全引号、修正逗号),提升结果稳定性。
二、BERT提取total_amount效果不佳的问题排查与优化
1. 常见问题排查
(1) 任务类型适配错误
若将任务设为文本分类(输入整段文本直接预测金额数值),BERT无法精确定位文本中的金额位置,效果必然差。正确做法是转为命名实体识别(NER)任务,标注total_amount对应的token标签。
(2) 数据标注问题
- 标注格式错误:未采用IOB/BIO标签体系(如
B-TOTAL表示金额起始token,I-TOTAL表示金额中间token,O表示其他); - 样本不平衡:多数样本的total_amount格式单一,或少数特殊格式样本(如带货币符号、千分位)标注不足;
- 标注错误:金额边界标注错误(如包含了“Total:”中的冒号,或漏标小数部分)。
(3) 训练参数与预处理问题
- 学习率过高/过低:BERT微调推荐学习率为
2e-5~5e-5,超出范围会导致模型不收敛或过拟合; - Token标签对齐错误:BERT的Tokenizer会拆分长词/数字为子词,若未将原文本的标签映射到子词级别,模型无法正确学习;
- 模型选择不当:使用了通用BERT(如
bert-base-uncased),未针对金融/数字文本优化的模型(如yiyanghkust/finbert-tone)。
2. 优化方案与修正代码
(1) 修正任务为NER(Token Classification)
示例标注数据(IOB格式)
输入文本:Total: $103.18
Token与标签对应:
| Token | Label |
|---|---|
| Total | O |
| : | O |
| $ | B-TOTAL |
| 103 | I-TOTAL |
| . | I-TOTAL |
| 18 | I-TOTAL |
核心训练代码
from transformers import BertTokenizer, BertForTokenClassification, TrainingArguments, Trainer from datasets import Dataset import torch # 示例数据集(已处理为token与标签对齐的格式) data = { "tokens": [["Amazon", ".", "com", "Invoice", "Total", ":", "$", "103", ".", "18"]], "labels": [[0, 0, 0, 0, 0, 0, 1, 2, 2, 2]] # 0=O,1=B-TOTAL,2=I-TOTAL } dataset = Dataset.from_dict(data) # 加载模型与Tokenizer tokenizer = BertTokenizer.from_pretrained('bert-base-cased') model = BertForTokenClassification.from_pretrained('bert-base-cased', num_labels=3) # 数据预处理:处理子词拆分后的标签对齐 def preprocess_function(examples): tokenized_inputs = tokenizer( examples["tokens"], is_split_into_words=True, truncation=True, max_length=512, padding='max_length' ) labels = [] for i, label in enumerate(examples["labels"]): word_ids = tokenized_inputs.word_ids(batch_index=i) previous_word_idx = None label_ids = [] for word_idx in word_ids: # 对padding token标注为-100(模型会忽略) if word_idx is None: label_ids.append(-100) # 同一单词的子词继承原标签,仅起始子词保留B标签,其余为I elif word_idx != previous_word_idx: label_ids.append(label[word_idx]) else: if label[word_idx] == 1: # 若原标签是B-TOTAL,子词转为I-TOTAL label_ids.append(2) else: label_ids.append(label[word_idx]) previous_word_idx = word_idx labels.append(label_ids) tokenized_inputs["labels"] = labels return tokenized_inputs tokenized_dataset = dataset.map(preprocess_function, batched=True) # 设置训练参数 training_args = TrainingArguments( output_dir='./total_amount_model', per_device_train_batch_size=8, num_train_epochs=4, learning_rate=3e-5, weight_decay=0.01, logging_dir='./logs', logging_steps=10, evaluation_strategy='epoch', save_strategy='epoch' ) # 自定义损失函数(可选,解决样本不平衡) def compute_metrics(p): predictions, labels = p predictions = torch.argmax(torch.tensor(predictions), dim=2) # 计算精确率、召回率、F1(针对TOTAL实体) from seqeval.metrics import precision_score, recall_score, f1_score true_labels = [[["O", "B-TOTAL", "I-TOTAL"][l] for l in label if l != -100] for label in labels] true_predictions = [[["O", "B-TOTAL", "I-TOTAL"][p] for (p, l) in zip(prediction, label) if l != -100] for prediction, label in zip(predictions, labels)] return { "precision": precision_score(true_labels, true_predictions), "recall": recall_score(true_labels, true_predictions), "f1": f1_score(true_labels, true_predictions) } # 初始化Trainer trainer = Trainer( model=model, args=training_args, train_dataset=tokenized_dataset, eval_dataset=tokenized_dataset, compute_metrics=compute_metrics ) # 训练 trainer.train() # 推理示例 def extract_total_amount(text): inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=512) outputs = model(**inputs) predictions = torch.argmax(outputs.logits, dim=2) tokens = tokenizer.convert_ids_to_tokens(inputs["input_ids"][0]) total_tokens = [] for token, pred in zip(tokens, predictions[0]): if pred in [1,2]: # B-TOTAL或I-TOTAL # 去掉子词前缀## if token.startswith("##"): total_tokens.append(token[2:]) else: total_tokens.append(token) total_str = "".join(total_tokens).replace("$", "").strip() try: return float(total_str) except: return None
(2) 其他优化建议
- 数据增强:对发票文本做随机打乱(金额位置不变)、添加不同货币符号、修改小数位数等,提升模型泛化性;
- 模型替换:使用金融领域预训练模型(如
yiyanghkust/finbert-tone),这类模型对数字、金融术语的理解更强; - 后处理优化:提取到金额字符串后,添加格式校验(如处理千分位逗号、货币符号),确保转换为正确数值。
内容的提问来源于stack exchange,提问作者Harshit Dubey
相关产品推荐
相关产品推荐

