Repository navigation
Expand file tree
/
Copy pathdata_utils.py
More file actions
121 lines (99 loc) · 4.32 KB
/
Copy pathdata_utils.py
File metadata and controls
121 lines (99 loc) · 4.32 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
import json
import os
import torch
from torch.utils.data import Dataset, DataLoader
from collections import Counter
import jieba
from tqdm import tqdm
class TranslationDataset(Dataset):
def __init__(self, file_path, src_tokenizer, tgt_tokenizer, max_len=128):
self.data = []
with open(file_path, 'r', encoding='utf-8') as f:
for line in f:
self.data.append(json.loads(line))
self.src_tokenizer = src_tokenizer
self.tgt_tokenizer = tgt_tokenizer
self.max_len = max_len
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
item = self.data[idx]
src_text = item['zh']
tgt_text = item['en']
src_ids = self.src_tokenizer.encode(src_text, add_special_tokens=True)
tgt_ids = self.tgt_tokenizer.encode(tgt_text, add_special_tokens=True)
# Truncate
if len(src_ids) > self.max_len:
src_ids = src_ids[:self.max_len-1] + [self.src_tokenizer.eos_token_id]
if len(tgt_ids) > self.max_len:
tgt_ids = tgt_ids[:self.max_len-1] + [self.tgt_tokenizer.eos_token_id]
return torch.tensor(src_ids), torch.tensor(tgt_ids)
def pad_collate_fn(batch, pad_idx):
src_list, tgt_list = [], []
for src, tgt in batch:
src_list.append(src)
tgt_list.append(tgt)
src_padded = torch.nn.utils.rnn.pad_sequence(src_list, batch_first=True, padding_value=pad_idx)
tgt_padded = torch.nn.utils.rnn.pad_sequence(tgt_list, batch_first=True, padding_value=pad_idx)
return src_padded, tgt_padded
class SimpleTokenizer:
def __init__(self, lang='en', vocab_size=None, min_freq=1):
self.lang = lang
self.min_freq = min_freq
self.vocab = {"<pad>": 0, "<unk>": 1, "<bos>": 2, "<eos>": 3}
self.id2token = {v: k for k, v in self.vocab.items()}
self.pad_token_id = 0
self.unk_token_id = 1
self.bos_token_id = 2
self.eos_token_id = 3
def build_vocab(self, texts):
counter = Counter()
for text in tqdm(texts, desc=f"Building vocab for {self.lang}"):
tokens = self.tokenize(text)
counter.update(tokens)
for token, freq in counter.most_common():
if freq >= self.min_freq:
if token not in self.vocab:
idx = len(self.vocab)
self.vocab[token] = idx
self.id2token[idx] = token
def tokenize(self, text):
if self.lang == 'zh':
return list(jieba.cut(text))
else:
return text.lower().replace('.', ' .').replace(',', ' ,').replace('?', ' ?').replace('!', ' !').split()
def encode(self, text, add_special_tokens=True):
tokens = self.tokenize(text)
ids = [self.vocab.get(t, self.unk_token_id) for t in tokens]
if add_special_tokens:
ids = [self.bos_token_id] + ids + [self.eos_token_id]
return ids
def decode(self, ids):
tokens = [self.id2token.get(i, "<unk>") for i in ids]
# Remove special tokens for clean text
clean_tokens = [t for t in tokens if t not in ["<pad>", "<bos>", "<eos>"]]
if self.lang == 'zh':
return "".join(clean_tokens)
else:
return " ".join(clean_tokens)
if __name__ == "__main__":
# Test
train_path = "/250010108/nlp/dataset/AP0004_Midterm&Final_translation_dataset_zh_en/train_10k.jsonl"
with open(train_path, 'r') as f:
lines = [json.loads(l) for l in f.readlines()]
zh_texts = [l['zh'] for l in lines]
en_texts = [l['en'] for l in lines]
src_tokenizer = SimpleTokenizer(lang='zh')
src_tokenizer.build_vocab(zh_texts)
tgt_tokenizer = SimpleTokenizer(lang='en')
tgt_tokenizer.build_vocab(en_texts)
print(f"ZH Vocab size: {len(src_tokenizer.vocab)}")
print(f"EN Vocab size: {len(tgt_tokenizer.vocab)}")
dataset = TranslationDataset(train_path, src_tokenizer, tgt_tokenizer)
dataloader = DataLoader(dataset, batch_size=4, collate_fn=lambda b: pad_collate_fn(b, src_tokenizer.pad_token_id))
for src, tgt in dataloader:
print("Source shape:", src.shape)
print("Target shape:", tgt.shape)
print("Decoded Source:", src_tokenizer.decode(src[0].tolist()))
print("Decoded Target:", tgt_tokenizer.decode(tgt[0].tolist()))
break