如何在HuggingFace NMT模型中跳过自定义术语表的分词与翻译?
方案解答
你的思路合理性判断
你提出的「正则匹配术语+特殊标记锁定译法」的方案完全合理,是无资源微调场景下实现术语硬约束翻译的最优轻量方案之一。核心逻辑是提前将术语替换为指定德语译法并包裹特殊标记,让模型只翻译未被标记的内容,同时通过标记避免模型篡改术语译法。
关键优化点与实现步骤
1. 术语匹配优先级处理
因为术语包含1-5元ngram,必须按最长ngram优先匹配,否则短术语会先匹配并破坏长术语的完整匹配(比如先匹配"custom"会导致"custom string"无法被识别)。
2. 特殊标记选择
建议使用模型陌生的、结构明确的组合标记(比如[[[...]]]或<<TERM>>...<</TERM>>),避免用单一"UKN"——单一标记容易被模型误翻译,而组合标记能明确告诉模型这部分是不可修改的固定内容。如果担心标记被tokenizer拆分,可以临时将标记加入tokenizer词汇表(无需微调模型)。
3. 具体实现代码
结合你的现有代码,修改后的完整流程如下:
import re from transformers import MBartForConditionalGeneration, AutoTokenizer # 1. 预处理术语表:按ngram长度从长到短排序,最长匹配优先 # 假设你的术语表是dict格式:term_dict = {"custom string": "desired string", ...} term_dict = {"custom string": "desired string", "AI model": "KI-Modell"} # 按术语长度(空格分割后的词数)倒序排序 sorted_terms = sorted(term_dict.keys(), key=lambda x: len(x.split()), reverse=True) # 生成正则模式:转义术语中的特殊字符,用|连接,确保最长匹配 pattern = re.compile(r'\b(' + '|'.join(re.escape(term) for term in sorted_terms) + r')\b', re.IGNORECASE) # 2. 加载模型和tokenizer model_path = "facebook/mbart-large-50-many-to-many-mmt" # 或opus-mt-en-de model = MBartForConditionalGeneration.from_pretrained(model_path) tokenizer = AutoTokenizer.from_pretrained(model_path) # 可选:添加自定义标记到tokenizer,避免被拆分 custom_tokens = ["[[[", "]]]"] tokenizer.add_tokens(custom_tokens) # 注意:无需微调模型,只是让tokenizer识别这些标记为单个token model.resize_token_embeddings(len(tokenizer)) # 3. 定义替换函数:匹配术语→替换为带标记的德语译法 def replace_terms(text): def replace_match(match): term = match.group(0) # 不区分大小写匹配术语表(如果需要严格区分则去掉lower()) lower_term = term.lower() # 找到对应的德语译法 de_term = next((v for k, v in term_dict.items() if k.lower() == lower_term), term) # 用标记包裹译法 return f"[[[{de_term}]]]" return pattern.sub(replace_match, text) # 4. 翻译流程 src_texts = ["longer sentence containing custom string etc.", "We use AI model for translation"] # 先替换术语 processed_texts = [replace_terms(text) for text in src_texts] # 适配tokenizer的正确翻译逻辑 tokenizer.src_lang = "en" inputs = tokenizer(processed_texts, return_tensors="pt", padding=True, truncation=True) translated_tokens = model.generate( **inputs, forced_bos_token_id=tokenizer.lang_code_to_id["de"], max_length=100 ) translated_texts = tokenizer.batch_decode(translated_tokens, skip_special_tokens=True) # 5. 移除标记 final_translations = [re.sub(r'\[\[\[(.*?)\]\]\]', r'\1', text) for text in translated_texts] print(final_translations) # 输出示例: # ["Längerer Satz mit desired string usw.", "Wir verwenden KI-Modell für die Übersetzung"]
注意事项
- 如果术语存在重叠(比如"custom"和"custom string"),最长匹配优先策略会优先匹配长术语,避免错误替换。
- 若发现模型仍修改标记内的内容,可以尝试更复杂的标记(比如
###TERM###...###END###),或确保标记被tokenizer识别为单个token。 - 对于大小写敏感的术语,可以去掉
re.IGNORECASE,并确保术语表的键与输入文本大小写完全一致。
内容的提问来源于stack exchange,提问作者Bharatiya
相关产品推荐
相关产品推荐

