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

Function build_model_input

modelzoo/din/train.py:438–489  ·  view source on GitHub ↗
(filename, batch_size, num_epochs)

Source from the content-addressed store, hash-verified

436
437# generate dataset pipline
438def build_model_input(filename, batch_size, num_epochs):
439 def parse_csv(value):
440 tf.logging.info('Parsing {}'.format(filename))
441 cate_defaults = [[' '] for i in range(0, 5)]
442 label_defaults = [[0]]
443 column_headers = TRAIN_DATA_COLUMNS
444 record_defaults = label_defaults + cate_defaults
445 columns = tf.io.decode_csv(value,
446 record_defaults=record_defaults,
447 field_delim='\t')
448 all_columns = collections.OrderedDict(zip(column_headers, columns))
449
450 labels = all_columns.pop(LABEL_COLUMN[0])
451 features = all_columns
452 return features, labels
453
454 def parse_parquet(value):
455 labels = value.pop(LABEL_COLUMN[0])
456 features = value
457 return features, labels
458
459 '''Work Queue Feature'''
460 if args.workqueue and not args.tf:
461 from tensorflow.python.ops.work_queue import WorkQueue
462 work_queue = WorkQueue([filename], num_epochs=num_epochs)
463 # For multiple files:
464 # work_queue = WorkQueue([filename, filename1,filename2,filename3])
465 files = work_queue.input_dataset()
466 else:
467 files = filename
468 # Extract lines from input files using the Dataset API.
469 if args.parquet_dataset and not args.tf:
470 from tensorflow.python.data.experimental.ops import parquet_dataset_ops
471 dataset = parquet_dataset_ops.ParquetDataset(files, batch_size=batch_size)
472 if args.parquet_dataset_shuffle:
473 dataset = dataset.shuffle(buffer_size=20000,
474 seed=args.seed) # set seed for reproducing
475 if not args.workqueue:
476 dataset = dataset.repeat(num_epochs)
477 dataset = dataset.map(parse_parquet,
478 num_parallel_calls=tf.data.experimental.AUTOTUNE)
479 else:
480 dataset = tf.data.TextLineDataset(files)
481 dataset = dataset.shuffle(buffer_size=20000,
482 seed=args.seed) # set seed for reproducing
483 if not args.workqueue:
484 dataset = dataset.repeat(num_epochs)
485 dataset = dataset.batch(batch_size)
486 dataset = dataset.map(parse_csv,
487 num_parallel_calls=tf.data.experimental.AUTOTUNE)
488 dataset = dataset.prefetch(2)
489 return dataset
490
491
492# generate feature columns

Callers 1

mainFunction · 0.70

Calls 7

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

Tested by

no test coverage detected