diff --git a/docs/usage.rst b/docs/usage.rst index edc69569..1a12a857 100644 --- a/docs/usage.rst +++ b/docs/usage.rst @@ -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 diff --git a/reagent/data/__init__.py b/reagent/data/__init__.py new file mode 100644 index 00000000..5be5087f --- /dev/null +++ b/reagent/data/__init__.py @@ -0,0 +1,2 @@ +#!/usr/bin/env python3 +# Copyright (c) Facebook, Inc. and its affiliates. All rights reserved. diff --git a/reagent/data/data_fetcher.py b/reagent/data/data_fetcher.py new file mode 100644 index 00000000..29e1db1a --- /dev/null +++ b/reagent/data/data_fetcher.py @@ -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() diff --git a/reagent/workflow/data_fetcher.py b/reagent/data/oss_data_fetcher.py similarity index 88% rename from reagent/workflow/data_fetcher.py rename to reagent/data/oss_data_fetcher.py index 306a5c86..da9d8ab7 100644 --- a/reagent/workflow/data_fetcher.py +++ b/reagent/data/oss_data_fetcher.py @@ -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) diff --git a/reagent/workflow/spark_utils.py b/reagent/data/spark_utils.py similarity index 100% rename from reagent/workflow/spark_utils.py rename to reagent/data/spark_utils.py diff --git a/reagent/model_managers/actor_critic_base.py b/reagent/model_managers/actor_critic_base.py index 7b52e23c..dd55d379 100644 --- a/reagent/model_managers/actor_critic_base.py +++ b/reagent/model_managers/actor_critic_base.py @@ -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, diff --git a/reagent/model_managers/discrete_dqn_base.py b/reagent/model_managers/discrete_dqn_base.py index 5a25a659..52400c60 100644 --- a/reagent/model_managers/discrete_dqn_base.py +++ b/reagent/model_managers/discrete_dqn_base.py @@ -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, diff --git a/reagent/model_managers/model_manager.py b/reagent/model_managers/model_manager.py index fcc6f5ee..c268178f 100644 --- a/reagent/model_managers/model_manager.py +++ b/reagent/model_managers/model_manager.py @@ -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 diff --git a/reagent/model_managers/parametric_dqn_base.py b/reagent/model_managers/parametric_dqn_base.py index 9c23f0b6..d4b2edd0 100644 --- a/reagent/model_managers/parametric_dqn_base.py +++ b/reagent/model_managers/parametric_dqn_base.py @@ -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() diff --git a/reagent/model_managers/policy_gradient/ppo.py b/reagent/model_managers/policy_gradient/ppo.py index 7467acd6..0e1a9422 100644 --- a/reagent/model_managers/policy_gradient/ppo.py +++ b/reagent/model_managers/policy_gradient/ppo.py @@ -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 diff --git a/reagent/model_managers/policy_gradient/reinforce.py b/reagent/model_managers/policy_gradient/reinforce.py index 70ba5280..2af53c5c 100644 --- a/reagent/model_managers/policy_gradient/reinforce.py +++ b/reagent/model_managers/policy_gradient/reinforce.py @@ -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 diff --git a/reagent/model_managers/slate_q_base.py b/reagent/model_managers/slate_q_base.py index 15752a7a..d71d7610 100644 --- a/reagent/model_managers/slate_q_base.py +++ b/reagent/model_managers/slate_q_base.py @@ -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") diff --git a/reagent/model_managers/world_model_base.py b/reagent/model_managers/world_model_base.py index f74d4955..2ac96efd 100644 --- a/reagent/model_managers/world_model_base.py +++ b/reagent/model_managers/world_model_base.py @@ -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() diff --git a/reagent/test/workflow/reagent_sql_test_base.py b/reagent/test/workflow/reagent_sql_test_base.py index a1f24250..1b20b01e 100644 --- a/reagent/test/workflow/reagent_sql_test_base.py +++ b/reagent/test/workflow/reagent_sql_test_base.py @@ -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 diff --git a/reagent/test/workflow/test_oss_workflows.py b/reagent/test/workflow/test_oss_workflows.py index 447514ce..ffe2274c 100644 --- a/reagent/test/workflow/test_oss_workflows.py +++ b/reagent/test/workflow/test_oss_workflows.py @@ -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, diff --git a/reagent/test/workflow/test_query_data.py b/reagent/test/workflow/test_query_data.py index ec8a183f..a1e256b5 100644 --- a/reagent/test/workflow/test_query_data.py +++ b/reagent/test/workflow/test_query_data.py @@ -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"], diff --git a/reagent/test/workflow/test_query_data_parametric.py b/reagent/test/workflow/test_query_data_parametric.py index 0c231dfd..0c8ddf4b 100644 --- a/reagent/test/workflow/test_query_data_parametric.py +++ b/reagent/test/workflow/test_query_data_parametric.py @@ -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, diff --git a/reagent/workflow/data/manual_data_module.py b/reagent/workflow/data/manual_data_module.py index 666bc47d..1b7d04bf 100644 --- a/reagent/workflow/data/manual_data_module.py +++ b/reagent/workflow/data/manual_data_module.py @@ -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 diff --git a/reagent/workflow/gym_batch_rl.py b/reagent/workflow/gym_batch_rl.py index f919906a..f5165198 100644 --- a/reagent/workflow/gym_batch_rl.py +++ b/reagent/workflow/gym_batch_rl.py @@ -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 diff --git a/reagent/workflow/identify_types_flow.py b/reagent/workflow/identify_types_flow.py index 92218b92..9e4566bd 100644 --- a/reagent/workflow/identify_types_flow.py +++ b/reagent/workflow/identify_types_flow.py @@ -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 diff --git a/reagent/workflow/training.py b/reagent/workflow/training.py index 486d3a47..209e6152 100644 --- a/reagent/workflow/training.py +++ b/reagent/workflow/training.py @@ -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") diff --git a/reagent/workflow/utils.py b/reagent/workflow/utils.py index 841ddd14..40345e9f 100644 --- a/reagent/workflow/utils.py +++ b/reagent/workflow/utils.py @@ -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