A class that handles one dataset
| 237 | yield [data["response"], data["context"]] |
| 238 | |
| 239 | class Dataset: |
| 240 | """ |
| 241 | A class that handles one dataset |
| 242 | """ |
| 243 | def __init__(self, filepath): |
| 244 | self.filepath = filepath |
| 245 | |
| 246 | def __iter__(self): |
| 247 | max_dataset_size = 10*1000*1000 #Cache small datasets in memory |
| 248 | dataset = [] |
| 249 | data_format = None |
| 250 | |
| 251 | while dataset is None or len(dataset) == 0: |
| 252 | with gzip.open(self.filepath, "rt") as fIn: |
| 253 | for line in fIn: |
| 254 | data = json.loads(line) |
| 255 | if isinstance(data, dict): |
| 256 | data = data['texts'] |
| 257 | |
| 258 | if data_format is None: |
| 259 | data_format = len(data) |
| 260 | |
| 261 | #Ensure that all entries are of the same 2/3 col format |
| 262 | assert len(data) == data_format |
| 263 | |
| 264 | if dataset is not None: |
| 265 | dataset.append(data) |
| 266 | if len(dataset) >= max_dataset_size: |
| 267 | dataset = None |
| 268 | |
| 269 | yield data |
| 270 | |
| 271 | # Data loaded. Now stream to the queue |
| 272 | # Shuffle for each epoch |
| 273 | while True: |
| 274 | random.shuffle(dataset) |
| 275 | for data in dataset: |
| 276 | yield data |
| 277 | |
| 278 | |
| 279 |