33import torch .utils .data
44from flgo .benchmark .toolkits .nlp .classification import GeneralCalculator
55from flgo .benchmark .base import FromDatasetPipe , FromDatasetGenerator
6- from torchtext .vocab import build_vocab_from_iterator
7- from torchtext .data .utils import get_tokenizer , ngrams_iterator
6+ from torchtext .data .utils import ngrams_iterator
87from torchtext .data .functional import to_map_style_dataset
98try :
109 import ujson as json
1110except :
1211 import json
1312from .config import train_data
14- try :
15- from .config import tokenizer
16- except :
17- tokenizer = None
18- try :
19- from .config import ngrams
20- except :
21- ngrams = 1
22-
23- def yield_tokens (data_iter , ngrams ):
24- for _ , text in data_iter :
25- yield ngrams_iterator (tokenizer (text ), ngrams )
26-
27- try :
28- from .config import vocab
29- except :
30- vocab = None
3113try :
3214 from .config import test_data
3315except :
@@ -37,30 +19,21 @@ def yield_tokens(data_iter, ngrams):
3719except :
3820 val_data = None
3921
40- if tokenizer is None : tokenizer = get_tokenizer ('basic_english' )
41- if vocab is None :
42- vocab = build_vocab_from_iterator (yield_tokens (train_data , ngrams ), specials = ["<unk>" ])
43- vocab .set_default_index (vocab ["<unk>" ])
44-
4522def collate_batch (batch ):
46- label_list , text_list , offsets = [], [], [ 0 ]
47- for (_label , _text ) in batch :
48- label_list .append (int ( _label ) - 1 )
49- processed_text = torch .tensor (vocab ( list ( ngrams_iterator ( tokenizer ( _text ), ngrams ))) , dtype = torch .int64 )
23+ label_list , text_list = [], []
24+ for (_text , _label ) in batch :
25+ label_list .append (_label )
26+ processed_text = torch .tensor (_text , dtype = torch .int64 )
5027 text_list .append (processed_text )
51- offsets .append (processed_text .size (0 ))
5228 label_list = torch .tensor (label_list , dtype = torch .int64 )
53- offsets = torch .tensor (offsets [:- 1 ]).cumsum (dim = 0 )
54- text_list = torch .cat (text_list )
55- return label_list , text_list , offsets
29+ return text_list , label_list
5630
5731class TaskGenerator (FromDatasetGenerator ):
5832 def __init__ (self ):
5933 super (TaskGenerator , self ).__init__ (benchmark = os .path .split (os .path .dirname (__file__ ))[- 1 ],
6034 train_data = train_data , val_data = val_data , test_data = test_data )
6135
6236 def prepare_data_for_partition (self ):
63- self .train_data = self .train_data .map (lambda x : (x [1 ], x [0 ]))
6437 return to_map_style_dataset (self .train_data )
6538
6639class TaskPipe (FromDatasetPipe ):
@@ -70,7 +43,7 @@ def __init__(self, task_path):
7043
7144 def save_task (self , generator ):
7245 client_names = self .gen_client_names (len (generator .local_datas ))
73- feddata = {'client_names' : client_names , }
46+ feddata = {'client_names' : client_names }
7447 for cid in range (len (client_names )): feddata [client_names [cid ]] = {'data' : generator .local_datas [cid ],}
7548 with open (os .path .join (self .task_path , 'data.json' ), 'w' ) as outf :
7649 json .dump (feddata , outf )
0 commit comments