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

Function build_model_input

modelzoo/dlrm/train.py:289–338  ·  view source on GitHub ↗
(filename, batch_size, num_epochs)

Source from the content-addressed store, hash-verified

287
288# generate dataset pipline
289def build_model_input(filename, batch_size, num_epochs):
290 def parse_csv(value):
291 tf.logging.info('Parsing {}'.format(filename))
292 cont_defaults = [[0.0] for i in range(1, 14)]
293 cate_defaults = [[' '] for i in range(1, 27)]
294 label_defaults = [[0]]
295 column_headers = TRAIN_DATA_COLUMNS
296 record_defaults = label_defaults + cont_defaults + cate_defaults
297 columns = tf.io.decode_csv(value, record_defaults=record_defaults)
298 all_columns = collections.OrderedDict(zip(column_headers, columns))
299 labels = all_columns.pop(LABEL_COLUMN[0])
300 features = all_columns
301 return features, labels
302
303 def parse_parquet(value):
304 tf.logging.info('Parsing {}'.format(filename))
305 labels = value.pop(LABEL_COLUMN[0])
306 features = value
307 return features, labels
308
309 '''Work Queue Feature'''
310 if args.workqueue and not args.tf:
311 from tensorflow.python.ops.work_queue import WorkQueue
312 work_queue = WorkQueue([filename], num_epochs=num_epochs)
313 # For multiple files:
314 # work_queue = WorkQueue([filename, filename1,filename2,filename3])
315 files = work_queue.input_dataset()
316 else:
317 files = filename
318
319 # Extract lines from input files using the Dataset API.
320 if args.parquet_dataset and not args.tf:
321 from tensorflow.python.data.experimental.ops import parquet_dataset_ops
322 dataset = parquet_dataset_ops.ParquetDataset(files, batch_size=batch_size)
323 if args.parquet_dataset_shuffle:
324 dataset = dataset.shuffle(buffer_size=20000,
325 seed=args.seed) # fix seed for reproducing
326 if not args.workqueue:
327 dataset = dataset.repeat(num_epochs)
328 dataset = dataset.map(parse_parquet, num_parallel_calls=28)
329 else:
330 dataset = tf.data.TextLineDataset(files)
331 dataset = dataset.shuffle(buffer_size=20000,
332 seed=args.seed) # fix seed for reproducing
333 if not args.workqueue:
334 dataset = dataset.repeat(num_epochs)
335 dataset = dataset.batch(batch_size)
336 dataset = dataset.map(parse_csv, num_parallel_calls=28)
337 dataset = dataset.prefetch(2)
338 return dataset
339
340
341# 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