MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / build_model_input

Function build_model_input

modelzoo/dien/train.py:604–669  ·  view source on GitHub ↗
(filename, neg_filename, batch_size, num_epochs)

Source from the content-addressed store, hash-verified

602
603# generate dataset pipline
604def build_model_input(filename, neg_filename, batch_size, num_epochs):
605 def parse_csv(value, neg_value):
606 tf.logging.info('Parsing {}'.format(filename))
607 cate_defaults = [[' '] for i in range(0, 5)]
608 label_defaults = [[0]]
609 column_headers = TRAIN_DATA_COLUMNS
610 record_defaults = label_defaults + cate_defaults
611 columns = tf.io.decode_csv(value,
612 record_defaults=record_defaults,
613 field_delim='\t')
614 neg_columns = tf.io.decode_csv(neg_value,
615 record_defaults=[[''], ['']],
616 field_delim='\t')
617 columns.extend(neg_columns)
618 all_columns = collections.OrderedDict(zip(column_headers, columns))
619
620 labels = all_columns.pop(LABEL_COLUMN[0])
621 features = all_columns
622 return features, labels
623
624 def parse_parquet(value, neg_value):
625 tf.logging.info('Parsing {}'.format(filename))
626 value.update(neg_value)
627 labels = value.pop(LABEL_COLUMN[0])
628 features = value
629 return features, labels
630
631 '''Work Queue Feature'''
632 if args.workqueue and not args.tf:
633 from tensorflow.python.ops.work_queue import WorkQueue
634 work_queue = WorkQueue([filename], num_epochs=num_epochs)
635 neg_work_queue = WorkQueue([neg_filename], num_epochs=num_epochs)
636 # For multiple files:
637 # work_queue = WorkQueue([filename, filename1,filename2,filename3])
638 files = work_queue.input_dataset()
639 neg_files = neg_work_queue.input_dataset()
640 else:
641 files = filename
642 neg_files = neg_filename
643 # Extract lines from input files using the Dataset API.
644 if args.parquet_dataset and not args.tf:
645 from tensorflow.python.data.experimental.ops import parquet_dataset_ops
646 dataset = parquet_dataset_ops.ParquetDataset(files, batch_size=batch_size)
647 dataset_neg_samples = parquet_dataset_ops.ParquetDataset(neg_files,
648 batch_size)
649 dataset = tf.data.Dataset.zip((dataset, dataset_neg_samples))
650 if args.parquet_dataset_shuffle:
651 dataset = dataset.shuffle(buffer_size=20000,
652 seed=args.seed) # set seed for reproducing
653 if not args.workqueue:
654 dataset = dataset.repeat(num_epochs)
655 dataset = dataset.map(parse_parquet,
656 num_parallel_calls=tf.data.experimental.AUTOTUNE)
657 else:
658 dataset = tf.data.TextLineDataset(files)
659 dataset_neg_samples = tf.data.TextLineDataset(neg_files)
660 dataset = tf.data.Dataset.zip((dataset, dataset_neg_samples))
661 dataset = dataset.shuffle(buffer_size=20000,

Callers 1

mainFunction · 0.70

Calls 8

input_datasetMethod · 0.95
WorkQueueClass · 0.90
zipMethod · 0.45
shuffleMethod · 0.45
repeatMethod · 0.45
mapMethod · 0.45
batchMethod · 0.45
prefetchMethod · 0.45

Tested by

no test coverage detected