如何基于'B'标签计数精准拆分NER BIO数据集的元组列表?
问题:BIO标注NER数据集拆分时B标签计数超标的问题
问题背景
我正在处理一个BIO标注的NER小数据集,需要手动拆分训练/验证/测试集。第一步要按'B'标签计数把数据拆成bin_1和bin_2,目标是bin_1的'B'总数为10,但当前代码拆分后总是超过这个数值。
数据集
data = [[('a', 'B'), ('b', 'I'), ('c', 'O'), ('d', 'B'), ('e', 'I'), ('f', 'O')], [('g', 'O'), ('h', 'O')], [('i', 'B'), ('j', 'I'), ('k', 'O')], [('l', 'B'), ('m', ''), ('n', 'B'), ('o', 'O')], [('p', 'O'), ('q', 'O'), ('r', 'O')], [('s', 'B'), ('t', 'O')], [('u', 'O'), ('v', 'B'), ('w', 'I'), ('x', 'O'), ('y', 'O')], [('z', 'B')], [('a', 'B'), ('b', 'I'), ('c', 'O')], [('d', 'O')], [('e', 'O'), ('f', 'O')], [('g', 'O'), ('h', 'B')], [('i', 'B'), ('j', 'I')], [('k', 'O')], [('l', 'O'), ('m', 'O'), ('n', 'O'), ('o', 'O')], [('p', 'O'), ('q', 'O'), ('r', 'O'), ('s', 'B'), ('t', 'O')], [('u', 'O'), ('v', 'B'), ('w', 'I'), ('x', 'O'), ('y', 'O'), ('z', 'B')]]
当前代码
import random # 原代码遗漏导入,此处补充 split = 0.7 d = [] total_B = 0 bin_1 = [] bin_2 = [] counter = 0 random.shuffle(data) for f in data: cnt = {} for _, label in f: if label in cnt: cnt[label] += 1 else: cnt[label] = 1 d.append(cnt) for f in d: total_B += f.get('B', 0) for f,g in zip(d, data): if f.get('B') is not None: if counter <= round(total_B * split): counter += f.get('B') bin_1.append(g) else: bin_2.append(g) print(f"Total count of 'B' in 'bin_1' should be: {round(total_B * split)}") print(f"Total count of 'B' in 'bin_1' is': {sum(1 for sublist in bin_1 for tuple_item in sublist if tuple_item[1] == 'B')}") print(f"Total count of 'B' in 'bin_2' is': {sum(1 for sublist in bin_2 for tuple_item in sublist if tuple_item[1] == 'B')}")
当前输出
Total count of 'B' in 'bin_1' should be: 10 Total count of 'B' in 'bin_1' is': 11 Total count of 'B' in 'bin_2' is': 3
bin_1, bin_2 >>> [[('a', 'B'), ('b', 'I'), ('c', 'O')], [('g', 'O'), ('h', 'B')], [('i', 'B'), ('j', 'I'), ('k', 'O')], [('u', 'O'), ('v', 'B'), ('w', 'I'), ('x', 'O'), ('y', 'O'), ('z', 'B')], [('s', 'B'), ('t', 'O')], [('l', 'B'), ('m', ''), ('n', 'B'), ('o', 'O')], [('a', 'B'), ('b', 'I'), ('c', 'O'), ('d', 'B'), ('e', 'I'), ('f', 'O')], [('i', 'B'), ('j', 'I')]], [[('u', 'O'), ('v', 'B'), ('w', 'I'), ('x', 'O'), ('y', 'O')], [('z', 'B')], [('p', 'O'), ('q', 'O'), ('r', 'O'), ('s', 'B'), ('t', 'O')]]
期望输出
Total count of 'B' in 'bin_1' should be: 10 Total count of 'B' in 'bin_1' is': 10 Total count of 'B' in 'bin_2' is': 4
问题原因
原代码的判断逻辑counter <= round(total_B * split)存在漏洞:当当前counter加上当前样本的B数会超过目标值时,仍然会将该样本加入bin_1,导致总数超标。比如目标为10,当counter为9时,遇到含2个B的样本,判断9 <=10成立,加入后counter变为11,直接超出目标。此外,原代码会忽略不含B的样本,导致这部分数据既不进入bin_1也不进入bin_2。
解决方案
修改判断逻辑,确保只有加入当前样本后counter不超过目标值时,才将样本放入bin_1;同时处理不含B的样本,避免数据丢失。
修正后的代码:
import random split = 0.7 bin_1 = [] bin_2 = [] counter = 0 random.shuffle(data) # 统计每个样本的B标签数量 sample_b_counts = [sum(1 for _, label in sample if label == 'B') for sample in data] total_B = sum(sample_b_counts) target = round(total_B * split) for b_count, sample in zip(sample_b_counts, data): if b_count == 0: # 不含B的样本直接加入bin_1(可根据需求调整分配逻辑) bin_1.append(sample) continue # 仅当加入后不超过目标时,才放入bin_1 if counter + b_count <= target: counter += b_count bin_1.append(sample) else: bin_2.append(sample) # 输出结果 print(f"Total count of 'B' in 'bin_1' should be: {target}") print(f"Total count of 'B' in 'bin_1' is': {sum(1 for sublist in bin_1 for tuple_item in sublist if tuple_item[1] == 'B')}") print(f"Total count of 'B' in 'bin_2' is': {sum(1 for sublist in bin_2 for tuple_item in sublist if tuple_item[1] == 'B')}")
说明
- 提前统计每个样本的B标签数量,简化后续逻辑,避免重复计算;
- 明确处理不含B的样本,避免数据丢失;
- 调整判断条件为
counter + b_count <= target,确保加入样本后不会超出目标值; - 代码结构更简洁,逻辑更清晰。
内容的提问来源于stack exchange,提问作者doine
相关产品推荐
相关产品推荐

