| 288 | # generate dataset pipline |
| 289 | def 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)) |