| 317 | |
| 318 | # generate dataset pipline |
| 319 | def 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 |
| 363 | def build_feature_columns(stock_tf, |