361 lines
13 KiB
C#
361 lines
13 KiB
C#
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; }
|
||
}
|
||
} |