CloudBuilder/CloudBuilder.AI.Service/Utility/ChineseNameExtractor.cs
owenchen e60118eb53 ow
2026-05-26 14:03:34 +08:00

361 lines
13 KiB
C#
Raw 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.

using CloudBuilder.Core.DatabaseAccessor.Entity;
using CloudBuilder.Core.Service;
using Microsoft.ML;
using Microsoft.ML.Data;
using Microsoft.ML.Transforms.Onnx;
using System;
using System.Collections.Generic;
using System.IO;
using System.Linq;
namespace CloudBuilder.AI.Service.Utility
{
/// <summary>
/// 基于 ONNX 模型的命名实体识别器,支持提取人名及人物角色(人名+职位)。
/// </summary>
public class ChineseNameExtractor
{
private BertChineseTokenizer tokenizer;
private OnnxScoringEstimator pipeline;
private PredictionEngine<NerInput, NerOutput> engine;
/// <summary>
/// 开启后会将每个 token 的预测标签及置信度输出到控制台,方便调试。
/// </summary>
public bool DebugMode { get; set; } = false;
public ChineseNameExtractor(IApplicationService service)
{
InitMLContext(service);
}
private void InitMLContext(IApplicationService service)
{
var mlContext = new MLContext();
string modelDir = service.Configuration["FileServiceSettings:OnnxDirectory"];
string modelPath = Path.Combine(modelDir, "model.onnx");
string vocabPath = Path.Combine(modelDir, "vocab.txt");
if (!File.Exists(modelPath) || !File.Exists(vocabPath))
throw new Exception("请确保 model.onnx 和 vocab.txt 在指定目录下。");
tokenizer = new BertChineseTokenizer(vocabPath);
var shapeDictionary = new Dictionary<string, int[]>
{
{ "input_ids", new[] { 1, 128 } },
{ "attention_mask", new[] { 1, 128 } },
{ "token_type_ids", new[] { 1, 128 } },
{ "logits", new[] { 1, 128, tokenizer.LabelCount } }
};
var emptyData = mlContext.Data.LoadFromEnumerable(new List<NerInput>());
pipeline = mlContext.Transforms.ApplyOnnxModel(
modelFile: modelPath,
inputColumnNames: new[] { "input_ids", "attention_mask", "token_type_ids" },
outputColumnNames: new[] { "logits" },
shapeDictionary: shapeDictionary,
gpuDeviceId: null,
fallbackToCpu: true
);
var transformer = pipeline.Fit(emptyData);
engine = mlContext.Model.CreatePredictionEngine<NerInput, NerOutput>(transformer);
}
/// <summary>
/// 提取纯人名(标签为 B-PER / I-PER 的连续实体)。
/// </summary>
public List<string> ExtractNames(string text)
{
return ExtractPersonEntities(text, onlyPureName: true);
}
/// <summary>
/// 提取人物角色(人名 + 紧接其后的职位,例如“岳堂主”、“张经理”等)。
/// </summary>
public List<string> ExtractPersonRoles(string text)
{
return ExtractPersonEntities(text, onlyPureName: false);
}
/// <summary>
/// 核心实体提取方法,支持合并人名与其后的职位标签。
/// </summary>
/// <param name="text">输入文本</param>
/// <param name="onlyPureName">true 只提取纯人名false 提取人名+职位复合实体</param>
private List<string> ExtractPersonEntities(string text, bool onlyPureName)
{
var tokenized = tokenizer.Tokenize(text);
var input = new NerInput
{
InputIds = tokenized.InputIds,
AttentionMask = tokenized.AttentionMask,
TokenTypeIds = tokenized.TokenTypeIds
};
var prediction = engine.Predict(input);
var idToLabel = tokenizer.IdToLabel;
int seqLen = 128;
int numLabels = idToLabel.Length;
var entities = new List<string>();
string currentEntity = "";
bool insideEntity = false;
bool isPersonEntity = false; // 当前实体是否以 B-PER 开始
for (int i = 0; i < seqLen; i++)
{
if (tokenized.AttentionMask[i] == 0) continue; // 跳过 padding
// 找到当前 token 的预测标签
int startIdx = i * numLabels;
float maxVal = float.MinValue;
int maxIdx = 0;
for (int j = 0; j < numLabels; j++)
{
float val = prediction.Logits[startIdx + j];
if (val > maxVal)
{
maxVal = val;
maxIdx = j;
}
}
string label = idToLabel[maxIdx];
string token = tokenized.Tokens[i];
if (DebugMode)
Console.WriteLine($"{token} -> {label} (置信度: {maxVal:F3})");
// 判断当前标签类型
bool isPerStart = label == "B-PER";
bool isPerInside = label == "I-PER";
bool isPositionStart = label == "B-POSITION";
bool isPositionInside = label == "I-POSITION";
// ---- 仅提取纯人名的逻辑 ----
if (onlyPureName)
{
if (isPerStart)
{
if (insideEntity && currentEntity.Length > 0) entities.Add(currentEntity);
currentEntity = token;
insideEntity = true;
isPersonEntity = true;
}
else if (isPerInside && insideEntity && isPersonEntity)
{
currentEntity += token;
}
else
{
if (insideEntity && currentEntity.Length > 0)
{
entities.Add(currentEntity);
currentEntity = "";
insideEntity = false;
isPersonEntity = false;
}
}
continue;
}
// ---- 提取人物角色(人名+职位)的逻辑 ----
if (isPerStart)
{
// 无论之前是否有未结束的实体,都结束之前的,并开始一个新的人名
if (insideEntity && currentEntity.Length > 0)
entities.Add(currentEntity);
currentEntity = token;
insideEntity = true;
isPersonEntity = true;
}
else if (isPositionStart)
{
if (insideEntity && isPersonEntity)
{
// 当前正在收集人名,且遇到 B-POSITION则合并职位
currentEntity += token;
// 保持 isPersonEntity = true 不变,这样后续的 I-POSITION 也能继续合并
}
else
{
// 不在人名内部,则视作新实体(提取单独的职位,如“堂主”)
if (insideEntity && currentEntity.Length > 0)
entities.Add(currentEntity);
currentEntity = token;
insideEntity = true;
isPersonEntity = false; // 标记为纯职位实体
}
}
else if ((isPerInside || isPositionInside) && insideEntity)
{
// 无论是 I-PER 还是 I-POSITION只要在实体内部都追加
currentEntity += token;
}
else
{
// 遇到其他标签O 或其他 B- 开头),结束当前实体
if (insideEntity && currentEntity.Length > 0)
{
entities.Add(currentEntity);
currentEntity = "";
insideEntity = false;
isPersonEntity = false;
}
// 注意:不处理非目标标签本身,因为它们不属于角色
}
}
// 收尾最后一个实体
if (insideEntity && currentEntity.Length > 0)
entities.Add(currentEntity);
return entities;
}
}
// ==================== 数据模型 ====================
public class NerInput
{
[VectorType(1, 128)]
[ColumnName("input_ids")]
public long[] InputIds { get; set; }
[VectorType(1, 128)]
[ColumnName("attention_mask")]
public long[] AttentionMask { get; set; }
[VectorType(1, 128)]
[ColumnName("token_type_ids")]
public long[] TokenTypeIds { get; set; }
}
public class NerOutput
{
[VectorType(1, 128, 29)] // 维度需与模型输出一致
[ColumnName("logits")]
public float[] Logits { get; set; }
}
// ==================== 分词器 ====================
public class BertChineseTokenizer
{
private readonly Dictionary<string, int> _tokenToId;
/// <summary>
/// 标签映射,务必与模型训练时的标签顺序一致。
/// </summary>
public readonly string[] IdToLabel;
public int LabelCount => IdToLabel.Length;
public BertChineseTokenizer(string vocabPath, int maxSeqLength = 128)
{
_tokenToId = new Dictionary<string, int>();
var lines = File.ReadAllLines(vocabPath);
for (int i = 0; i < lines.Length; i++)
_tokenToId[lines[i]] = i;
// 请根据实际模型输出维度调整此数组!
// 下面的顺序仅为示例,必须与模型 logits 第2维一一对应。
IdToLabel = new string[] {
"O", // 0
"B-ORG", // 1
"I-ORG", // 2
"B-PER", // 3
"I-PER", // 4
"B-TIME", // 5
"I-TIME", // 6
"B-LOC", // 7
"I-LOC", // 8
"B-POSITION", // 9
"I-POSITION", // 10
"B-COMPANY", // 11
"I-COMPANY", // 12
"B-GAME", // 13
"I-GAME", // 14
"B-GOVERNMENT", // 15
"I-GOVERNMENT", // 16
"B-SCENE", // 17
"I-SCENE", // 18
"B-SUBJECT", // 19
"I-SUBJECT", // 20
"B-CREATION", // 21
"I-CREATION", // 22
"B-FOOD", // 23
"I-FOOD", // 24
"B-MOVIE", // 25
"I-MOVIE", // 26
"[CLS]", // 27 (如果模型没有这些特殊标签,请删除)
"[SEP]" // 28
};
}
public TokenizedResult Tokenize(string text)
{
const int maxLen = 128;
var inputIds = new long[maxLen];
var attentionMask = new long[maxLen];
var tokenTypeIds = new long[maxLen];
var tokens = new string[maxLen];
var chars = text.ToCharArray();
int charIndex = 0;
for (int i = 0; i < maxLen; i++)
{
if (i == 0)
{
inputIds[i] = 101;
tokens[i] = "[CLS]";
}
else if (i == chars.Length + 1)
{
inputIds[i] = 102;
tokens[i] = "[SEP]";
}
else if (i > chars.Length + 1)
{
inputIds[i] = 0;
tokens[i] = "[PAD]";
}
else
{
string c = chars[charIndex].ToString();
charIndex++;
if (_tokenToId.TryGetValue(c, out int id))
{
inputIds[i] = id;
tokens[i] = c;
}
else
{
inputIds[i] = _tokenToId["[UNK]"];
tokens[i] = "[UNK]";
}
}
attentionMask[i] = (i <= chars.Length + 1) ? 1 : 0;
tokenTypeIds[i] = 0;
}
return new TokenizedResult
{
InputIds = inputIds,
AttentionMask = attentionMask,
TokenTypeIds = tokenTypeIds,
Tokens = tokens
};
}
}
public class TokenizedResult
{
public long[] InputIds { get; set; }
public long[] AttentionMask { get; set; }
public long[] TokenTypeIds { get; set; }
public string[] Tokens { get; set; }
}
}