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

Function build_model_input

modelzoo/bst/train.py:371–424  ·  view source on GitHub ↗
(filename, batch_size, num_epochs)

Source from the content-addressed store, hash-verified

369
370# generate dataset pipline
371def 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

Callers 1

mainFunction · 0.70

Calls 7

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

Tested by

no test coverage detected