Move data fetcher out of workflow (#445)

Summary: Pull Request resolved: https://github.com/facebookresearch/ReAgent/pull/445

Reviewed By: kaiwenw

Differential Revision: D27303639

fbshipit-source-id: 1c8f105a90aa929c8fecae12aa3191a0a8ed0008
This commit is contained in:
Jason Gauci
2021-04-07 16:18:54 -07:00
committed by Facebook GitHub Bot
parent 9cd616f33d
commit c133d7b012
22 changed files with 120 additions and 62 deletions
+1
View File
@@ -231,6 +231,7 @@ To train the model, we first save our Spark table to Parquet format, and use `Pe
input_table_spec=input_table_spec, # description of Spark table
sample_range=train_sample_range, # what percentage of data to use for training
reward_options=reward_options, # config to calculate rewards
data_fetcher=data_fetcher, # Controller for fetching data
)
# train_dataset now points to a Parquet
+2
View File
@@ -0,0 +1,2 @@
#!/usr/bin/env python3
# Copyright (c) Facebook, Inc. and its affiliates. All rights reserved.
+23
View File
@@ -0,0 +1,23 @@
#!/usr/bin/env python3
import logging
from typing import List, Optional, Tuple
from reagent.workflow.types import Dataset, TableSpec
logger = logging.getLogger(__name__)
class DataFetcher:
def query_data(
self,
input_table_spec: TableSpec,
discrete_action: bool,
actions: Optional[List[str]] = None,
include_possible_actions=True,
custom_reward_expression: Optional[str] = None,
sample_range: Optional[Tuple[float, float]] = None,
multi_steps: Optional[int] = None,
gamma: Optional[float] = None,
) -> Dataset:
raise NotImplementedError()
@@ -14,9 +14,9 @@ from pyspark.sql.types import (
StructField,
StructType,
)
from .spark_utils import get_spark_session, get_table_url
from .types import Dataset, TableSpec
from reagent.data.data_fetcher import DataFetcher
from reagent.data.spark_utils import get_spark_session, get_table_url
from reagent.workflow.types import Dataset, TableSpec
logger = logging.getLogger(__name__)
@@ -428,51 +428,56 @@ def upload_as_parquet(df) -> Dataset:
return Dataset(parquet_url=parquet_url)
def query_data(
input_table_spec: TableSpec,
discrete_action: bool,
actions: Optional[List[str]] = None,
include_possible_actions=True,
custom_reward_expression: Optional[str] = None,
sample_range: Optional[Tuple[float, float]] = None,
multi_steps: Optional[int] = None,
gamma: Optional[float] = None,
) -> Dataset:
"""Perform reward calculation, hashing mdp + subsampling and
other preprocessing such as sparse2dense.
"""
sqlCtx = get_spark_session()
df = sqlCtx.sql(f"SELECT * FROM {input_table_spec.table_name}")
df = set_reward_col_as_reward(
df,
custom_reward_expression=custom_reward_expression,
multi_steps=multi_steps,
gamma=gamma,
)
df = hash_mdp_id_and_subsample(df, sample_range=sample_range)
df = misc_column_preprocessing(df, multi_steps=multi_steps)
df = state_and_metrics_sparse2dense(
df,
states=infer_states_names(df, multi_steps),
metrics=infer_metrics_names(df, multi_steps),
multi_steps=multi_steps,
)
if discrete_action:
assert include_possible_actions
assert actions is not None, "in discrete case, actions must be given."
df = discrete_action_preprocessing(df, actions=actions, multi_steps=multi_steps)
else:
actions = infer_action_names(df, multi_steps)
df = parametric_action_preprocessing(
class OssDataFetcher(DataFetcher):
def query_data(
self,
input_table_spec: TableSpec,
discrete_action: bool,
actions: Optional[List[str]] = None,
include_possible_actions=True,
custom_reward_expression: Optional[str] = None,
sample_range: Optional[Tuple[float, float]] = None,
multi_steps: Optional[int] = None,
gamma: Optional[float] = None,
) -> Dataset:
"""Perform reward calculation, hashing mdp + subsampling and
other preprocessing such as sparse2dense.
"""
sqlCtx = get_spark_session()
# pyre-ignore
df = sqlCtx.sql(f"SELECT * FROM {input_table_spec.table_name}")
df = set_reward_col_as_reward(
df,
actions=actions,
custom_reward_expression=custom_reward_expression,
multi_steps=multi_steps,
gamma=gamma,
)
df = hash_mdp_id_and_subsample(df, sample_range=sample_range)
df = misc_column_preprocessing(df, multi_steps=multi_steps)
df = state_and_metrics_sparse2dense(
df,
states=infer_states_names(df, multi_steps),
metrics=infer_metrics_names(df, multi_steps),
multi_steps=multi_steps,
)
if discrete_action:
assert include_possible_actions
assert actions is not None, "in discrete case, actions must be given."
df = discrete_action_preprocessing(
df, actions=actions, multi_steps=multi_steps
)
else:
actions = infer_action_names(df, multi_steps)
df = parametric_action_preprocessing(
df,
actions=actions,
multi_steps=multi_steps,
include_possible_actions=include_possible_actions,
)
df = select_relevant_columns(
df,
discrete_action=discrete_action,
include_possible_actions=include_possible_actions,
)
df = select_relevant_columns(
df,
discrete_action=discrete_action,
include_possible_actions=include_possible_actions,
)
return upload_as_parquet(df)
return upload_as_parquet(df)
+3 -2
View File
@@ -13,6 +13,7 @@ from reagent.core.parameters import (
NormalizationData,
NormalizationKey,
)
from reagent.data.data_fetcher import DataFetcher
from reagent.evaluation.evaluator import get_metrics_to_score
from reagent.gym.policies.policy import Policy
from reagent.gym.policies.predictor_policies import create_predictor_policy_from_model
@@ -26,7 +27,6 @@ from reagent.preprocessing.batch_preprocessor import (
from reagent.preprocessing.normalization import get_feature_config
from reagent.preprocessing.types import InputColumn
from reagent.workflow.data import ReAgentDataModule
from reagent.workflow.data_fetcher import query_data
from reagent.workflow.identify_types_flow import identify_normalization_parameters
from reagent.workflow.reporters.actor_critic_reporter import ActorCriticReporter
from reagent.workflow.types import (
@@ -204,9 +204,10 @@ class ActorCriticBase(ModelManager):
input_table_spec: TableSpec,
sample_range: Optional[Tuple[float, float]],
reward_options: RewardOptions,
data_fetcher: DataFetcher,
) -> Dataset:
logger.info("Starting query")
return query_data(
return data_fetcher.query_data(
input_table_spec=input_table_spec,
discrete_action=False,
include_possible_actions=False,
+4 -2
View File
@@ -10,6 +10,7 @@ from reagent.core.parameters import (
NormalizationData,
NormalizationKey,
)
from reagent.data.data_fetcher import DataFetcher
from reagent.evaluation.evaluator import get_metrics_to_score
from reagent.gym.policies.policy import Policy
from reagent.gym.policies.predictor_policies import create_predictor_policy_from_model
@@ -28,7 +29,6 @@ from reagent.preprocessing.preprocessor import Preprocessor
from reagent.preprocessing.types import InputColumn
from reagent.workflow.data import ReAgentDataModule
from reagent.workflow.data.manual_data_module import ManualDataModule
from reagent.workflow.data_fetcher import query_data
from reagent.workflow.identify_types_flow import identify_normalization_parameters
from reagent.workflow.reporters.discrete_dqn_reporter import DiscreteDQNReporter
from reagent.workflow.types import (
@@ -110,6 +110,7 @@ class DiscreteDQNBase(ModelManager):
input_table_spec: TableSpec,
sample_range: Optional[Tuple[float, float]],
reward_options: RewardOptions,
data_fetcher: DataFetcher,
) -> Dataset:
raise RuntimeError
@@ -227,8 +228,9 @@ class DiscreteDqnDataModule(ManualDataModule):
input_table_spec: TableSpec,
sample_range: Optional[Tuple[float, float]],
reward_options: RewardOptions,
data_fetcher: DataFetcher,
) -> Dataset:
return query_data(
return data_fetcher.query_data(
input_table_spec=input_table_spec,
discrete_action=True,
actions=self.model_manager.action_names,
+2
View File
@@ -9,6 +9,7 @@ import torch
from reagent.core.dataclasses import dataclass
from reagent.core.parameters import NormalizationData
from reagent.core.registry_meta import RegistryMeta
from reagent.data.data_fetcher import DataFetcher
from reagent.training import Trainer
from reagent.workflow.data import ReAgentDataModule
from reagent.workflow.types import (
@@ -151,6 +152,7 @@ class ModelManager(metaclass=RegistryMeta):
input_table_spec: TableSpec,
sample_range: Optional[Tuple[float, float]],
reward_options: RewardOptions,
data_fetcher: DataFetcher,
) -> Dataset:
"""
DEPRECATED: Implement get_data_module() instead
@@ -10,6 +10,7 @@ from reagent.core.parameters import (
NormalizationData,
NormalizationKey,
)
from reagent.data.data_fetcher import DataFetcher
from reagent.evaluation.evaluator import get_metrics_to_score
from reagent.gym.policies.policy import Policy
from reagent.gym.policies.predictor_policies import create_predictor_policy_from_model
@@ -150,6 +151,7 @@ class ParametricDQNBase(ModelManager):
input_table_spec: TableSpec,
sample_range: Optional[Tuple[float, float]],
reward_options: RewardOptions,
data_fetcher: DataFetcher,
) -> Dataset:
raise NotImplementedError()
@@ -9,6 +9,7 @@ from reagent.core.dataclasses import dataclass, field
from reagent.core.parameters import NormalizationData
from reagent.core.parameters import NormalizationKey
from reagent.core.parameters import param_hash
from reagent.data.data_fetcher import DataFetcher
from reagent.gym.policies.policy import Policy
from reagent.gym.policies.predictor_policies import create_predictor_policy_from_model
from reagent.gym.policies.samplers.discrete_sampler import SoftmaxActionSampler
@@ -122,6 +123,7 @@ class PPO(ModelManager):
input_table_spec: TableSpec,
sample_range: Optional[Tuple[float, float]],
reward_options: RewardOptions,
data_fetcher: DataFetcher,
) -> Dataset:
raise NotImplementedError
@@ -9,6 +9,7 @@ from reagent.core.dataclasses import dataclass, field
from reagent.core.parameters import NormalizationData
from reagent.core.parameters import NormalizationKey
from reagent.core.parameters import param_hash
from reagent.data.data_fetcher import DataFetcher
from reagent.gym.policies.policy import Policy
from reagent.gym.policies.predictor_policies import create_predictor_policy_from_model
from reagent.gym.policies.samplers.discrete_sampler import SoftmaxActionSampler
@@ -124,6 +125,7 @@ class Reinforce(ModelManager):
input_table_spec: TableSpec,
sample_range: Optional[Tuple[float, float]],
reward_options: RewardOptions,
data_fetcher: DataFetcher,
) -> Dataset:
raise NotImplementedError
+2
View File
@@ -5,6 +5,7 @@ from typing import Dict, List, Optional, Tuple
import reagent.core.types as rlt
from reagent.core.dataclasses import dataclass
from reagent.core.parameters import NormalizationData, NormalizationKey
from reagent.data.data_fetcher import DataFetcher
from reagent.gym.policies.policy import Policy
from reagent.gym.policies.predictor_policies import create_predictor_policy_from_model
from reagent.gym.policies.samplers.top_k_sampler import TopKSampler
@@ -140,6 +141,7 @@ class SlateQBase(ModelManager):
input_table_spec: TableSpec,
sample_range: Optional[Tuple[float, float]],
reward_options: RewardOptions,
data_fetcher: DataFetcher,
) -> Dataset:
raise NotImplementedError("Write for OSS")
@@ -4,6 +4,7 @@ from typing import Dict, List, Optional, Tuple
from reagent.core.dataclasses import dataclass
from reagent.core.parameters import NormalizationData, NormalizationKey
from reagent.data.data_fetcher import DataFetcher
from reagent.gym.policies.policy import Policy
from reagent.model_managers.model_manager import ModelManager
from reagent.preprocessing.batch_preprocessor import BatchPreprocessor
@@ -51,6 +52,7 @@ class WorldModelBase(ModelManager):
input_table_spec: TableSpec,
sample_range: Optional[Tuple[float, float]],
reward_options: RewardOptions,
data_fetcher: DataFetcher,
) -> Dataset:
raise NotImplementedError()
@@ -9,9 +9,7 @@ import shutil
import numpy as np
import torch
from pyspark import SparkConf
# pyre-fixme[21]: Could not find module `reagent.workflow.spark_utils`.
from reagent.workflow.spark_utils import DEFAULT_SPARK_CONFIG
from reagent.data.spark_utils import DEFAULT_SPARK_CONFIG
# pyre-fixme[21]: Could not find `sparktestingbase`.
from sparktestingbase.sqltestcase import SQLTestCase
+2 -1
View File
@@ -95,7 +95,8 @@ class TestOSSWorkflows(HorizonTestBase):
)
mock_normalization = mock_cartpole_normalization()
with patch(
f"{DISCRETE_DQN_BASE}.query_data", return_value=mock_dataset
"reagent.data.oss_data_fetcher.OssDataFetcher.query_data",
return_value=mock_dataset,
), patch(
f"{DISCRETE_DQN_BASE}.identify_normalization_parameters",
return_value=mock_normalization,
+3 -2
View File
@@ -9,11 +9,11 @@ import pytest
# pyre-ignore
from pyspark.sql.functions import asc # @manual=//python/wheel/pyspark:pyspark
from reagent.data.oss_data_fetcher import OssDataFetcher
from reagent.test.test_data.ex_mdps import generate_discrete_mdp_pandas_df
# pyre-ignore
from reagent.test.workflow.reagent_sql_test_base import ReagentSQLTestBase
from reagent.workflow.data_fetcher import query_data
from reagent.workflow.types import Dataset, TableSpec
@@ -47,7 +47,8 @@ class TestQueryData(ReagentSQLTestBase):
self, custom_reward_expression=None, gamma=None, multi_steps=None
):
ts = TableSpec(table_name=self.table_name)
dataset: Dataset = query_data(
df = OssDataFetcher()
dataset: Dataset = df.query_data(
input_table_spec=ts,
discrete_action=True,
actions=["L", "R", "U", "D"],
@@ -9,11 +9,11 @@ import pytest
# pyre-fixme[21]: Could not find `pyspark`.
from pyspark.sql.functions import asc
from reagent.data.oss_data_fetcher import OssDataFetcher
from reagent.test.test_data.ex_mdps import generate_parametric_mdp_pandas_df
# pyre-fixme[21]: Could not find `workflow`.
from reagent.test.workflow.reagent_sql_test_base import ReagentSQLTestBase
from reagent.workflow.data_fetcher import query_data
from reagent.workflow.types import Dataset, TableSpec
logger = logging.getLogger(__name__)
@@ -46,7 +46,8 @@ class TestQueryDataParametric(ReagentSQLTestBase):
self, custom_reward_expression=None, gamma=None, multi_steps=None
):
ts = TableSpec(table_name=self.table_name)
dataset: Dataset = query_data(
df = OssDataFetcher()
dataset: Dataset = df.query_data(
input_table_spec=ts,
discrete_action=False,
include_possible_actions=False,
@@ -20,6 +20,8 @@ except ModuleNotFoundError:
from reagent.core.parameters import NormalizationData
from reagent.data.data_fetcher import DataFetcher
from reagent.data.oss_data_fetcher import OssDataFetcher
from reagent.preprocessing.batch_preprocessor import (
BatchPreprocessor,
)
@@ -108,6 +110,8 @@ class ManualDataModule(ReAgentDataModule):
key = "normalization_data_map"
data_fetcher = OssDataFetcher()
normalization_data_map = (
self.run_feature_identification(self.input_table_spec)
if key not in self.saved_setup_data
@@ -121,6 +125,7 @@ class ManualDataModule(ReAgentDataModule):
input_table_spec=self.input_table_spec,
sample_range=sample_range_output.train_sample_range,
reward_options=self.reward_options,
data_fetcher=data_fetcher,
)
eval_dataset = None
if calc_cpe_in_training:
@@ -128,6 +133,7 @@ class ManualDataModule(ReAgentDataModule):
input_table_spec=self.input_table_spec,
sample_range=sample_range_output.eval_sample_range,
reward_options=self.reward_options,
data_fetcher=data_fetcher,
)
return self._pickle_setup_data(
@@ -228,6 +234,7 @@ class ManualDataModule(ReAgentDataModule):
input_table_spec: TableSpec,
sample_range: Optional[Tuple[float, float]],
reward_options: RewardOptions,
data_fetcher: DataFetcher,
) -> Dataset:
"""
Massage input table into the format expected by the trainer
+1 -1
View File
@@ -10,6 +10,7 @@ import gym
import numpy as np
import pandas as pd
import torch
from reagent.data.spark_utils import call_spark_class, get_spark_session
from reagent.gym.agents.agent import Agent
from reagent.gym.envs import Gym
from reagent.gym.policies.predictor_policies import create_predictor_policy_from_model
@@ -20,7 +21,6 @@ from reagent.publishers.union import FileSystemPublisher, ModelPublisher__Union
from reagent.replay_memory.circular_replay_buffer import ReplayBuffer
from reagent.replay_memory.utils import replay_buffer_to_pre_timeline_df
from .spark_utils import call_spark_class, get_spark_session
from .types import TableSpec
+1 -1
View File
@@ -8,12 +8,12 @@ import reagent.core.types as rlt
# pyre-fixme[21]: Could not find `pyspark`.
# pyre-fixme[21]: Could not find `pyspark`.
from pyspark.sql.functions import col, collect_list, explode
from reagent.data.spark_utils import get_spark_session
from reagent.preprocessing.normalization import (
NormalizationParameters,
get_feature_norm_metadata,
)
from .spark_utils import get_spark_session
from .types import PreprocessingOptions, TableSpec
+4
View File
@@ -8,6 +8,7 @@ from typing import Dict, Optional
import torch
from reagent.core.parameters import NormalizationData
from reagent.core.tensorboardX import summary_writer_context
from reagent.data.oss_data_fetcher import OssDataFetcher
from reagent.model_managers.model_manager import ModelManager
from reagent.model_managers.union import ModelManager__Union
from reagent.publishers.union import ModelPublisher__Union
@@ -138,6 +139,7 @@ def query_and_train(
train_dataset = None
eval_dataset = None
data_fetcher = OssDataFetcher()
if normalization_data_map is not None:
calc_cpe_in_training = manager.should_generate_eval_dataset
sample_range_output = get_sample_range(input_table_spec, calc_cpe_in_training)
@@ -145,6 +147,7 @@ def query_and_train(
input_table_spec=input_table_spec,
sample_range=sample_range_output.train_sample_range,
reward_options=reward_options,
data_fetcher=data_fetcher,
)
eval_dataset = None
if calc_cpe_in_training:
@@ -152,6 +155,7 @@ def query_and_train(
input_table_spec=input_table_spec,
sample_range=sample_range_output.eval_sample_range,
reward_options=reward_options,
data_fetcher=data_fetcher,
)
logger.info("Starting training")
+1 -1
View File
@@ -14,10 +14,10 @@ from petastorm import make_batch_reader
# pyre-fixme[21]: Could not find module `petastorm.pytorch`.
from petastorm.pytorch import DataLoader, decimal_friendly_collate
from pytorch_lightning.loggers import TensorBoardLogger
from reagent.data.spark_utils import get_spark_session
from reagent.preprocessing.batch_preprocessor import BatchPreprocessor
from reagent.training import StoppingEpochCallback
from .spark_utils import get_spark_session
from .types import Dataset, ReaderOptions, ResourceOptions