| 56 | |
| 57 | |
| 58 | def 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 | |
| 84 | def build_feature_cols(): |