| 369 | |
| 370 | # generate dataset pipline |
| 371 | def build_model_input(filename, batch_size, num_epochs): |
| 372 | def parse_csv(value): |
| 373 | tf.logging.info('Parsing {}'.format(filename)) |
| 374 | string_defaults = [[' '] for i in range(1, 19)] |
| 375 | label_defaults = [[0], [0]] |
| 376 | column_headers = INPUT_COLUMN |
| 377 | record_defaults = label_defaults + string_defaults |
| 378 | columns = tf.io.decode_csv(value, record_defaults=record_defaults) |
| 379 | all_columns = collections.OrderedDict(zip(column_headers, columns)) |
| 380 | labels = all_columns.pop(LABEL_COLUMN[0]) |
| 381 | all_columns.pop(BUY_COLUMN[0]) |
| 382 | features = all_columns |
| 383 | return features, labels |
| 384 | |
| 385 | def parse_parquet(value): |
| 386 | tf.logging.info('Parsing {}'.format(filename)) |
| 387 | labels = value.pop(LABEL_COLUMN[0]) |
| 388 | features = value |
| 389 | return features, labels |
| 390 | |
| 391 | '''Work Queue Feature''' |
| 392 | if args.workqueue and not args.tf: |
| 393 | from tensorflow.python.ops.work_queue import WorkQueue |
| 394 | work_queue = WorkQueue([filename], num_epochs=num_epochs) |
| 395 | # For multiple files: |
| 396 | # work_queue = WorkQueue([filename, filename1,filename2,filename3]) |
| 397 | files = work_queue.input_dataset() |
| 398 | else: |
| 399 | files = filename |
| 400 | # Extract lines from input files using the Dataset API. |
| 401 | if args.parquet_dataset and not args.tf: |
| 402 | from tensorflow.python.data.experimental.ops import parquet_dataset_ops |
| 403 | dataset = parquet_dataset_ops.ParquetDataset(files, |
| 404 | fields = PARQUET_INPUT_COLUMN, |
| 405 | batch_size=batch_size) |
| 406 | if args.parquet_dataset_shuffle: |
| 407 | dataset = dataset.shuffle(buffer_size=20000, |
| 408 | seed=args.seed) # fix seed for reproducing |
| 409 | if not args.workqueue: |
| 410 | dataset = dataset.repeat(num_epochs) |
| 411 | dataset = dataset.prefetch(batch_size) |
| 412 | dataset = dataset.map(parse_parquet, num_parallel_calls=28) |
| 413 | else: |
| 414 | dataset = tf.data.TextLineDataset(files) |
| 415 | dataset = dataset.shuffle(buffer_size=20000, |
| 416 | seed=args.seed) # fix seed for reproducing |
| 417 | if not args.workqueue: |
| 418 | dataset = dataset.repeat(num_epochs) |
| 419 | dataset = dataset.prefetch(batch_size) |
| 420 | dataset = dataset.batch(batch_size) |
| 421 | dataset = dataset.map(parse_csv, num_parallel_calls=28) |
| 422 | |
| 423 | dataset = dataset.prefetch(1) |
| 424 | return dataset |
| 425 | |
| 426 | |
| 427 | # generate feature columns |