Repository navigation
Expand file tree
/
Copy pathtest_ray_document_deduplicator.py
More file actions
254 lines (223 loc) · 8.76 KB
/
Copy pathtest_ray_document_deduplicator.py
File metadata and controls
254 lines (223 loc) · 8.76 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
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
import unittest
from unittest.mock import patch
from data_juicer.core.data import NestedDataset as Dataset
from data_juicer.ops.deduplicator.ray_document_deduplicator import \
RayDocumentDeduplicator
from data_juicer.utils.unittest_utils import DataJuicerTestCaseBase, TEST_TAG
class RayDocumentDeduplicatorTest(DataJuicerTestCaseBase):
def _run_ray_cross_block_dedup(self, samples, op):
dataset = self._build_ray_cross_block_dataset(samples)
dataset.process([op])
return dataset.data.take_all()
def _build_ray_cross_block_dataset(self, samples):
import ray
from data_juicer.core.data.ray_dataset import RayDataset
from data_juicer.utils.constant import Fields
ds_list = [{Fields.stats: {}, **sample} for sample in samples]
return RayDataset(
ray.data.from_items(ds_list, override_num_blocks=len(ds_list)),
cfg={'auto_op_parallelism': False},
auto_op_parallelism=False,
)
def _run_doc_dedup(self, dataset: Dataset, target_list, op):
import ray
from data_juicer.core.data.ray_dataset import RayDataset
dataset = RayDataset(
ray.data.from_items(dataset.to_list()),
cfg={'auto_op_parallelism': False},
auto_op_parallelism=False,
)
dataset.process([op])
res_list = [
{op.text_key: sample[op.text_key]}
for sample in dataset.data.take_all()
]
res_list.sort(key=lambda x: x['text'])
target_list.sort(key=lambda x: x['text'])
self.assertEqual(res_list, target_list)
@TEST_TAG("ray")
def test_english_deduplication(self):
ds_list = [
{
'text': 'Today is Sunday and it\'s a happy day!'
},
{
'text': 'Do you need a cup of coffee?'
},
{
'text': 'Today is sunday and it\'s a happy day!'
},
{
'text':
'This paper proposed a novel method on LLM pretraining.'
},
{
'text':
'This paper proposed a novel method on LLM pretraining.'
},
]
tgt_list = [{
'text': 'Today is Sunday and it\'s a happy day!'
}, {
'text': 'Do you need a cup of coffee?'
}, {
'text': 'Today is sunday and it\'s a happy day!'
}, {
'text':
'This paper proposed a novel method on LLM pretraining.'
}]
dataset = self.generate_dataset(ds_list)
op = RayDocumentDeduplicator(
lowercase=False,
ignore_non_character=False,
dedup_set_num=1,
batch_size=1,
num_proc=2,
auto_op_parallelism=False,
)
self._run_doc_dedup(dataset, tgt_list, op)
@TEST_TAG("ray")
def test_chinese_deduplication(self):
ds_list = [
{
'text': '你好,请问你是谁'
},
{
'text': '欢迎来到阿里巴巴!'
},
{
'text':
'第九届会议\n2003年7月28日至8月8日\n牙买加金斯敦\n为来自发展中国家的法'
'律和技术委员会以及财务委员会成员\n参加委员会会议支付费用的方式\n1.'
},
{
'text':
'第九届会议\n2003年7月28日至8月8日\n牙买加金斯敦\n为来自发展中国家的法'
'律和技术委员会以及财务委员会成员\n参加委员会会议支付费用的方式\n1.'
},
{
'text':
'第九届会议\n时间:2003年7月28日至8月8日\n牙买加金斯敦\n为来自发展中国家的法'
'律和技术委员会以及财务委员会成员\n参加委员会会议支付费用的方式\n1.'
},
]
tgt_list = [
{
'text': '你好,请问你是谁'
},
{
'text': '欢迎来到阿里巴巴!'
},
{
'text':
'第九届会议\n2003年7月28日至8月8日\n牙买加金斯敦\n为来自发展中国家的法'
'律和技术委员会以及财务委员会成员\n参加委员会会议支付费用的方式\n1.'
},
{
'text':
'第九届会议\n时间:2003年7月28日至8月8日\n牙买加金斯敦\n为来自发展中国家的法'
'律和技术委员会以及财务委员会成员\n参加委员会会议支付费用的方式\n1.'
},
]
dataset = self.generate_dataset(ds_list)
op = RayDocumentDeduplicator(
lowercase=False,
ignore_non_character=False,
dedup_set_num=1,
batch_size=1,
num_proc=2,
auto_op_parallelism=False,
)
self._run_doc_dedup(dataset, tgt_list, op)
@TEST_TAG("ray")
def test_ray_actor_backend_deduplicates_across_blocks(self):
op = RayDocumentDeduplicator(
lowercase=False,
ignore_non_character=False,
dedup_set_num=1,
batch_size=1,
num_proc=4,
auto_op_parallelism=False,
)
res_list = self._run_ray_cross_block_dedup(
[{'text': 'duplicate across ray blocks'} for _ in range(8)],
op,
)
self.assertEqual(len(res_list), 1)
self.assertEqual(res_list[0]['text'], 'duplicate across ray blocks')
@TEST_TAG("ray")
def test_ray_actor_execution_mode_still_shares_dedup_sets(self):
op = RayDocumentDeduplicator(
lowercase=False,
ignore_non_character=False,
dedup_set_num=1,
batch_size=1,
num_proc=4,
auto_op_parallelism=False,
ray_execution_mode='actor',
)
res_list = self._run_ray_cross_block_dedup(
[{'text': 'duplicate with actor execution mode'} for _ in range(8)],
op,
)
self.assertEqual(len(res_list), 1)
self.assertEqual(res_list[0]['text'], 'duplicate with actor execution mode')
@TEST_TAG("ray")
def test_ray_basic_deduplicator_subclasses_share_dedup_sets(self):
from data_juicer.ops.deduplicator.ray_image_deduplicator import RayImageDeduplicator
from data_juicer.ops.deduplicator.ray_video_deduplicator import RayVideoDeduplicator
cases = [
(RayImageDeduplicator, {'images': []}),
(RayVideoDeduplicator, {'videos': []}),
]
for op_cls, sample in cases:
with self.subTest(op_cls=op_cls.__name__):
op = op_cls(
dedup_set_num=1,
batch_size=1,
num_proc=4,
auto_op_parallelism=False,
)
res_list = self._run_ray_cross_block_dedup([sample for _ in range(8)], op)
self.assertEqual(len(res_list), 1)
@TEST_TAG("ray")
def test_repeated_execution_keeps_materialized_dedup_result(self):
op = RayDocumentDeduplicator(
lowercase=False,
ignore_non_character=False,
dedup_set_num=1,
batch_size=1,
num_proc=4,
auto_op_parallelism=False,
)
dataset = self._build_ray_cross_block_dataset([{
'text': 'duplicate across repeated executions',
} for _ in range(8)])
dataset.process([op])
self.assertEqual(dataset.data.count(), 1)
res_list = dataset.data.take_all()
self.assertEqual(len(res_list), 1)
self.assertEqual(res_list[0]['text'], 'duplicate across repeated executions')
@TEST_TAG("ray")
def test_stats_export_does_not_consume_dedup_state_before_filter(self):
def materializing_write_json(dataset, *args, **kwargs):
return dataset.count()
with patch('ray.data.Dataset.write_json', materializing_write_json):
op = RayDocumentDeduplicator(
lowercase=False,
ignore_non_character=False,
dedup_set_num=1,
batch_size=1,
num_proc=4,
auto_op_parallelism=False,
stats_export_path='mock_stats_export_path',
)
dataset = self._build_ray_cross_block_dataset([{
'text': 'duplicate with stats export',
} for _ in range(8)])
dataset.process([op])
res_list = dataset.data.take_all()
self.assertEqual(len(res_list), 1)
self.assertEqual(res_list[0]['text'], 'duplicate with stats export')
if __name__ == '__main__':
unittest.main()