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

Function generate_input_data

modelzoo/features/runtime/deepfm/train.py:58–81  ·  view source on GitHub ↗
(filename, batch_size, num_epochs)

Source from the content-addressed store, hash-verified

56
57
58def generate_input_data(filename, batch_size, num_epochs):
59 def parse_csv(value):
60 tf.logging.info('Parsing {}'.format(filename))
61 cont_defaults = [[0.0] for i in range(1, 14)]
62 cate_defaults = [[" "] for i in range(1, 27)]
63 label_defaults = [[0]]
64 column_headers = TRAIN_DATA_COLUMNS
65 record_defaults = label_defaults + cont_defaults + cate_defaults
66 columns = tf.io.decode_csv(value, record_defaults=record_defaults)
67 all_columns = collections.OrderedDict(zip(column_headers, columns))
68 labels = all_columns.pop(LABEL_COLUMN[0])
69 features = all_columns
70 return features, labels
71
72 # Extract lines from input files using the Dataset API.
73 dataset = tf.data.TextLineDataset(filename)
74 dataset = dataset.shuffle(buffer_size=20000,
75 seed=2021) # set seed for reproducing
76 dataset = dataset.repeat(num_epochs)
77 dataset = dataset.prefetch(batch_size)
78 dataset = dataset.batch(batch_size)
79 dataset = dataset.map(parse_csv, num_parallel_calls=28)
80 dataset = dataset.prefetch(1)
81 return dataset
82
83
84def build_feature_cols():

Callers 1

mainFunction · 0.70

Calls 5

shuffleMethod · 0.45
repeatMethod · 0.45
prefetchMethod · 0.45
batchMethod · 0.45
mapMethod · 0.45

Tested by

no test coverage detected