| 325 | |
| 326 | # generate dataset pipline |
| 327 | def build_model_input(filename, batch_size, num_epochs): |
| 328 | def parse_csv(value): |
| 329 | tf.logging.info('Parsing {}'.format(filename)) |
| 330 | cont_defaults = [[0.0] for i in range(1, 14)] |
| 331 | cate_defaults = [[' '] for i in range(1, 27)] |
| 332 | label_defaults = [[0]] |
| 333 | column_headers = TRAIN_DATA_COLUMNS |
| 334 | record_defaults = label_defaults + cont_defaults + cate_defaults |
| 335 | columns = tf.io.decode_csv(value, record_defaults=record_defaults) |
| 336 | all_columns = collections.OrderedDict(zip(column_headers, columns)) |
| 337 | labels = all_columns.pop(LABEL_COLUMN[0]) |
| 338 | features = all_columns |
| 339 | return features, labels |
| 340 | |
| 341 | def parse_parquet(value): |
| 342 | tf.logging.info('Parsing {}'.format(filename)) |
| 343 | labels = value.pop(LABEL_COLUMN[0]) |
| 344 | features = value |
| 345 | return features, labels |
| 346 | |
| 347 | '''Work Queue Feature''' |
| 348 | if args.workqueue and not args.tf: |
| 349 | from tensorflow.python.ops.work_queue import WorkQueue |
| 350 | work_queue = WorkQueue([filename], num_epochs=num_epochs) |
| 351 | # For multiple files: |
| 352 | # work_queue = WorkQueue([filename, filename1,filename2,filename3]) |
| 353 | files = work_queue.input_dataset() |
| 354 | else: |
| 355 | files = filename |
| 356 | # Extract lines from input files using the Dataset API. |
| 357 | if args.parquet_dataset and not args.tf: |
| 358 | from tensorflow.python.data.experimental.ops import parquet_dataset_ops |
| 359 | dataset = parquet_dataset_ops.ParquetDataset( |
| 360 | files, batch_size=batch_size) |
| 361 | if args.parquet_dataset_shuffle: |
| 362 | dataset = dataset.shuffle(buffer_size=20000, |
| 363 | seed=args.seed) # fix seed for reproducing |
| 364 | if not args.workqueue: |
| 365 | dataset = dataset.repeat(num_epochs) |
| 366 | dataset = dataset.map(parse_parquet, num_parallel_calls=28) |
| 367 | else: |
| 368 | dataset = tf.data.TextLineDataset(files) |
| 369 | dataset = dataset.shuffle(buffer_size=20000, |
| 370 | seed=args.seed) # fix seed for reproducing |
| 371 | if not args.workqueue: |
| 372 | dataset = dataset.repeat(num_epochs) |
| 373 | dataset = dataset.batch(batch_size) |
| 374 | dataset = dataset.map(parse_csv, num_parallel_calls=28) |
| 375 | return dataset |
| 376 | |
| 377 | |
| 378 | def build_feature_columns(): |