mirror of
https://github.com/facebookresearch/ReAgent.git
synced 2026-06-16 12:44:41 +00:00
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:
committed by
Facebook GitHub Bot
parent
9cd616f33d
commit
c133d7b012
@@ -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
|
||||
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) Facebook, Inc. and its affiliates. All rights reserved.
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user