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

Function build_model_input

modelzoo/dcn/train.py:360–408  ·  view source on GitHub ↗
(filename, batch_size, num_epochs)

Source from the content-addressed store, hash-verified

358
359# generate dataset pipline
360def build_model_input(filename, batch_size, num_epochs):
361 def parse_csv(value):
362 tf.logging.info('Parsing {}'.format(filename))
363 cont_defaults = [[0.0] for i in range(1, 14)]
364 cate_defaults = [[' '] for i in range(1, 27)]
365 label_defaults = [[0]]
366 column_headers = TRAIN_DATA_COLUMNS
367 record_defaults = label_defaults + cont_defaults + cate_defaults
368 columns = tf.io.decode_csv(value, record_defaults=record_defaults)
369 all_columns = collections.OrderedDict(zip(column_headers, columns))
370 labels = all_columns.pop(LABEL_COLUMN[0])
371 features = all_columns
372 return features, labels
373
374 def parse_parquet(value):
375 tf.logging.info('Parsing {}'.format(filename))
376 labels = value.pop(LABEL_COLUMN[0])
377 features = value
378 return features, labels
379
380 '''Work Queue Feature'''
381 if args.workqueue and not args.tf:
382 from tensorflow.python.ops.work_queue import WorkQueue
383 work_queue = WorkQueue([filename], num_epochs=num_epochs)
384 # For multiple files:
385 # work_queue = WorkQueue([filename, filename1,filename2,filename3])
386 files = work_queue.input_dataset()
387 else:
388 files = filename
389 # Extract lines from input files using the Dataset API.
390 if args.parquet_dataset and not args.tf:
391 from tensorflow.python.data.experimental.ops import parquet_dataset_ops
392 dataset = parquet_dataset_ops.ParquetDataset(files, batch_size=batch_size)
393 if args.parquet_dataset_shuffle:
394 dataset = dataset.shuffle(buffer_size=20000,
395 seed=args.seed) # fix seed for reproducing
396 if not args.workqueue:
397 dataset = dataset.repeat(num_epochs)
398 dataset = dataset.map(parse_parquet, num_parallel_calls=28)
399 else:
400 dataset = tf.data.TextLineDataset(files)
401 dataset = dataset.shuffle(buffer_size=20000,
402 seed=args.seed) # fix seed for reproducing
403 if not args.workqueue:
404 dataset = dataset.repeat(num_epochs)
405 dataset = dataset.batch(batch_size)
406 dataset = dataset.map(parse_csv, num_parallel_calls=28)
407 dataset = dataset.prefetch(2)
408 return dataset
409
410
411# 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