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

Function build_model_input

modelzoo/esmm/train.py:319–360  ·  view source on GitHub ↗
(filename, batch_size, num_epochs, seed, stock_tf, workqueue=None)

Source from the content-addressed store, hash-verified

317
318# generate dataset pipline
319def build_model_input(filename, batch_size, num_epochs, seed, stock_tf, workqueue=None):
320 def parse_csv(value):
321 tf.logging.info('Parsing {}'.format(filename))
322 columns = tf.io.decode_csv(value, record_defaults=defaults)
323 all_columns = collections.OrderedDict(zip(headers, columns))
324 labels = [all_columns.pop(LABEL_COLUMNS[0]), all_columns.pop(LABEL_COLUMNS[1])]
325 label = tf.multiply(labels[0], labels[1])
326 features = all_columns
327 return features, label
328
329 def parse_parquet(value):
330 tf.logging.info('Parsing {}'.format(filename))
331 labels = [value.pop(LABEL_COLUMNS[0]), value.pop(LABEL_COLUMNS[1])]
332 label = tf.multiply(labels[0], labels[1])
333 features = value
334 return features, label
335
336 '''Work Queue Feature'''
337 if not stock_tf and workqueue:
338 from tensorflow.python.ops.work_queue import WorkQueue
339 work_queue = WorkQueue([filename], num_epochs=num_epochs)
340 files = work_queue.input_dataset()
341 else:
342 files = filename
343 if args.parquet_dataset and not args.tf:
344 from tensorflow.python.data.experimental.ops import parquet_dataset_ops
345 dataset = parquet_dataset_ops.ParquetDataset(files, batch_size=batch_size)
346 if args.parquet_dataset_shuffle:
347 dataset = dataset.shuffle(buffer_size=10000,
348 seed=seed) # fix seed for reproducing
349 if not args.workqueue:
350 dataset = dataset.repeat(num_epochs)
351 dataset = dataset.map(parse_parquet, num_parallel_calls=tf.data.experimental.AUTOTUNE)
352 else:
353 dataset = tf.data.TextLineDataset(files)
354 dataset = dataset.shuffle(buffer_size=10000, seed=seed)
355 if not args.workqueue:
356 dataset = dataset.repeat(num_epochs)
357 dataset = dataset.batch(batch_size)
358 dataset = dataset.map(parse_csv, num_parallel_calls=tf.data.experimental.AUTOTUNE)
359 dataset = dataset.prefetch(2)
360 return dataset
361
362# generate feature columns
363def build_feature_columns(stock_tf,

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