mirror of
https://github.com/facebookresearch/ReAgent.git
synced 2026-06-16 12:44:41 +00:00
data-fetcher that preprocesses data and stores as Petastorm dataset (#223)
Summary: Converted data_fetcher.py’s query_data to PySpark and added sparse2dense logic. Since we’re using Parquet (which shall be loaded with Petastorm), instead of returning HiveDataSetClass, I’m returning Dataset containing url to parquet. Pull Request resolved: https://github.com/facebookresearch/ReAgent/pull/223 Test Plan: Imported from GitHub, without a `Test Plan:` line. Run `tox`. Reviewed By: kittipatv Differential Revision: D21036014 Pulled By: kaiwenw fbshipit-source-id: 46f1cef7731db365c70dce4831b6d7adca39dce1
This commit is contained in:
committed by
Facebook GitHub Bot
parent
d9bbcfafc2
commit
eae3c8a149
@@ -0,0 +1,40 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) Facebook, Inc. and its affiliates. All rights reserved.
|
||||
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from pyspark import SparkConf
|
||||
from sparktestingbase.sqltestcase import SQLTestCase
|
||||
|
||||
|
||||
SEED = 42
|
||||
|
||||
|
||||
class ReagentSQLTestBase(SQLTestCase):
|
||||
def getConf(self):
|
||||
conf = SparkConf()
|
||||
# set shuffle partitions to a low number, e.g. <= cores * 2 to speed
|
||||
# things up, otherwise the tests will use the default 200 partitions
|
||||
# and it will take a lot more time to complete
|
||||
conf.set("spark.sql.shuffle.partitions", "12")
|
||||
conf.set("spark.port.maxRetries", "30")
|
||||
return conf
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
assert not os.path.isdir("metastore_db"), "metastore_db already exists"
|
||||
|
||||
torch.manual_seed(SEED)
|
||||
np.random.seed(SEED)
|
||||
logging.basicConfig()
|
||||
|
||||
def tearDown(self):
|
||||
super().tearDown()
|
||||
|
||||
# removes Derby from last runs
|
||||
if os.path.isdir("metastore_db"):
|
||||
shutil.rmtree("metastore_db")
|
||||
@@ -1,22 +1,30 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) Facebook, Inc. and its affiliates. All rights reserved.
|
||||
|
||||
import logging
|
||||
import unittest
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
from reagent.preprocessing.identify_types import CONTINUOUS
|
||||
from reagent.test.base.horizon_test_base import HorizonTestBase
|
||||
from reagent.workflow.identify_types_flow import (
|
||||
get_spark_session,
|
||||
identify_normalization_parameters,
|
||||
)
|
||||
from reagent.test.workflow.reagent_sql_test_base import ReagentSQLTestBase
|
||||
from reagent.workflow.identify_types_flow import identify_normalization_parameters
|
||||
from reagent.workflow.types import PreprocessingOptions, TableSpec
|
||||
|
||||
|
||||
class TestPreprocessing(HorizonTestBase):
|
||||
def test_preprocessing(self):
|
||||
spark = get_spark_session()
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SEED = 42
|
||||
NUM_ROWS = 10000
|
||||
|
||||
|
||||
class TestPreprocessing(ReagentSQLTestBase):
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
logging.getLogger(__name__).setLevel(logging.INFO)
|
||||
|
||||
@pytest.mark.serial
|
||||
def test_preprocessing(self):
|
||||
distributions = {}
|
||||
distributions["0"] = {"mean": 0, "stddev": 1, "size": (5,)}
|
||||
distributions["1"] = {"mean": 4, "stddev": 3, "size": (3,)}
|
||||
@@ -29,26 +37,25 @@ class TestPreprocessing(HorizonTestBase):
|
||||
for k, info in distributions.items()
|
||||
}
|
||||
|
||||
np.random.seed(42)
|
||||
data = [(i, get_random_feature()) for i in range(100000)]
|
||||
df = spark.sparkContext.parallelize(data).toDF(["i", "states"])
|
||||
data = [(i, get_random_feature()) for i in range(NUM_ROWS)]
|
||||
df = self.sc.parallelize(data).toDF(["i", "states"])
|
||||
df.show()
|
||||
|
||||
table_name = "test_table"
|
||||
df.createOrReplaceTempView(table_name)
|
||||
|
||||
num_samples = 10000
|
||||
num_samples = NUM_ROWS / 10
|
||||
preprocessing_options = PreprocessingOptions(num_samples=num_samples)
|
||||
|
||||
table_spec = TableSpec(table_name=table_name)
|
||||
|
||||
normalization_params = identify_normalization_parameters(
|
||||
table_spec, "states", preprocessing_options, seed=42
|
||||
table_spec, "states", preprocessing_options, seed=SEED
|
||||
)
|
||||
|
||||
print(normalization_params)
|
||||
logger.info(normalization_params)
|
||||
for k, info in distributions.items():
|
||||
print(
|
||||
logger.info(
|
||||
f"Expect {k} to be normal with mean {info['mean']}, stddev {info['stddev']}"
|
||||
)
|
||||
assert normalization_params[k].feature_type == CONTINUOUS
|
||||
@@ -58,9 +65,7 @@ class TestPreprocessing(HorizonTestBase):
|
||||
assert abs(
|
||||
normalization_params[k].stddev - info["stddev"] < 0.2
|
||||
), f"{normalization_params[k].stddev} not close to {info['stddev']}"
|
||||
print("Everything seems fine.")
|
||||
|
||||
spark.stop()
|
||||
logger.info("identify_normalization_parameters seems fine.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -0,0 +1,329 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) Facebook, Inc. and its affiliates. All rights reserved.
|
||||
|
||||
import logging
|
||||
import tempfile
|
||||
import unittest
|
||||
from os.path import abspath
|
||||
|
||||
import numpy as np
|
||||
import pandas
|
||||
import pytest
|
||||
from pyspark.sql.functions import asc
|
||||
from reagent.test.environment.environment import MultiStepSamples
|
||||
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__)
|
||||
|
||||
|
||||
def generate_data_discrete(sqlCtx, multi_step: bool, table_name: str):
|
||||
# Simulate the following MDP:
|
||||
# state: 0, action: 7 ('L'), reward: 0,
|
||||
# state: 1, action: 8 ('R'), reward: 1,
|
||||
# state: 4, action: 9 ('U'), reward: 4,
|
||||
# state: 5, action: 10 ('D'), reward: 5,
|
||||
# state: 6 (terminal)
|
||||
actions = ["L", "R", "U", "D"]
|
||||
possible_actions = [["L", "R"], ["R", "U"], ["U", "D"], ["D"]]
|
||||
|
||||
# assume multi_step=2
|
||||
if multi_step:
|
||||
rewards = [[0, 1], [1, 4], [4, 5], [5]]
|
||||
metrics = [
|
||||
[{"reward": 0}, {"reward": 1}],
|
||||
[{"reward": 1}, {"reward": 4}],
|
||||
[{"reward": 4}, {"reward": 5}],
|
||||
[{"reward": 5}],
|
||||
]
|
||||
next_states = [[{1: 1}, {4: 1}], [{4: 1}, {5: 1}], [{5: 1}, {6: 1}], [{6: 1}]]
|
||||
next_actions = [["R", "U"], ["U", "D"], ["D", ""], [""]]
|
||||
possible_next_actions = [
|
||||
[["R", "U"], ["U", "D"]],
|
||||
[["U", "D"], ["D"]],
|
||||
[["D"], [""]],
|
||||
[[""]],
|
||||
]
|
||||
terminals = [[0, 0], [0, 0], [0, 1], [1]]
|
||||
time_diffs = [[1, 1], [1, 1], [1, 1], [1]]
|
||||
else:
|
||||
rewards = [[0], [1], [4], [5]]
|
||||
metrics = [{"reward": 0}, {"reward": 1}, {"reward": 4}, {"reward": 5}] # noqa
|
||||
next_states = [[{1: 1}], [{4: 1}], [{5: 1}], [{6: 1}]]
|
||||
next_actions = [["R"], ["U"], ["D"], [""]]
|
||||
possible_next_actions = [[["R", "U"]], [["U", "D"]], [["D"]], [[""]]]
|
||||
terminals = [[0], [0], [0], [1]]
|
||||
time_diffs = [1, 3, 1, 1] # noqa
|
||||
|
||||
samples = MultiStepSamples(
|
||||
mdp_ids=["0", "0", "0", "0"],
|
||||
sequence_numbers=[0, 1, 4, 5],
|
||||
sequence_number_ordinals=[1, 2, 3, 4],
|
||||
states=[{0: 1}, {1: 1}, {4: 1}, {5: 1}],
|
||||
actions=actions,
|
||||
action_probabilities=[0.3, 0.4, 0.5, 0.6],
|
||||
rewards=rewards,
|
||||
possible_actions=possible_actions,
|
||||
next_states=next_states,
|
||||
next_actions=next_actions,
|
||||
terminals=terminals,
|
||||
possible_next_actions=possible_next_actions,
|
||||
)
|
||||
if not multi_step:
|
||||
samples = samples.to_single_step()
|
||||
|
||||
next_state_features = samples.next_states
|
||||
possible_next_actions = samples.possible_next_actions
|
||||
next_actions = samples.next_actions
|
||||
|
||||
df = pandas.DataFrame(
|
||||
{
|
||||
"mdp_id": samples.mdp_ids,
|
||||
"sequence_number": samples.sequence_numbers,
|
||||
"sequence_number_ordinal": samples.sequence_number_ordinals,
|
||||
"state_features": samples.states,
|
||||
"action": samples.actions,
|
||||
"action_probability": samples.action_probabilities,
|
||||
"reward": samples.rewards,
|
||||
"next_state_features": next_state_features,
|
||||
"next_action": next_actions,
|
||||
"time_diff": time_diffs,
|
||||
"possible_actions": samples.possible_actions,
|
||||
"possible_next_actions": possible_next_actions,
|
||||
"metrics": metrics,
|
||||
}
|
||||
)
|
||||
df = sqlCtx.createDataFrame(df)
|
||||
logger.info("Created dataframe")
|
||||
df.show()
|
||||
df.createOrReplaceTempView(table_name)
|
||||
|
||||
|
||||
def assertEq(series_a, arr_b):
|
||||
arr_a = np.array(series_a.tolist())
|
||||
np.testing.assert_equal(arr_a, arr_b)
|
||||
|
||||
|
||||
def assertAllClose(series_a, arr_b):
|
||||
arr_a = np.array(series_a.tolist())
|
||||
np.testing.assert_allclose(arr_a, arr_b)
|
||||
|
||||
|
||||
def assertEqWithPresence(series_a, presence, arr_b):
|
||||
arr_a = np.array(series_a.tolist())
|
||||
present_a = arr_a[presence]
|
||||
present_b = arr_b[presence]
|
||||
np.testing.assert_equal(present_a, present_b)
|
||||
|
||||
|
||||
def verify_single_step_except_rewards(df):
|
||||
""" expects a pandas dataframe """
|
||||
assertEq(df["sequence_number"], pandas.Series([1, 2, 3, 4]))
|
||||
|
||||
state_features_presence = np.array(
|
||||
[
|
||||
[True, False, False, False, False],
|
||||
[False, True, False, False, False],
|
||||
[False, False, True, False, False],
|
||||
[False, False, False, True, False],
|
||||
],
|
||||
dtype="bool",
|
||||
)
|
||||
assertEq(df["state_features_presence"], state_features_presence)
|
||||
state_features = np.array(
|
||||
[
|
||||
[1.0, 0.0, 0.0, 0.0, 0.0],
|
||||
[0.0, 1.0, 0.0, 0.0, 0.0],
|
||||
[0.0, 0.0, 1.0, 0.0, 0.0],
|
||||
[0.0, 0.0, 0.0, 1.0, 0.0],
|
||||
],
|
||||
dtype="float32",
|
||||
)
|
||||
assertEqWithPresence(df["state_features"], state_features_presence, state_features)
|
||||
|
||||
assertEq(df["action"], np.array([0, 1, 2, 3]))
|
||||
assertEq(df["action_probability"], np.array([0.3, 0.4, 0.5, 0.6], dtype="float32"))
|
||||
assertEq(df["not_terminal"], np.array([1, 1, 1, 0], dtype="bool"))
|
||||
next_state_features_presence = np.array(
|
||||
[
|
||||
[False, True, False, False, False],
|
||||
[False, False, True, False, False],
|
||||
[False, False, False, True, False],
|
||||
[False, False, False, False, True],
|
||||
],
|
||||
dtype="bool",
|
||||
)
|
||||
assertEq(df["next_state_features_presence"], next_state_features_presence)
|
||||
next_state_features = np.array(
|
||||
[
|
||||
[0.0, 1.0, 0.0, 0.0, 0.0],
|
||||
[0.0, 0.0, 1.0, 0.0, 0.0],
|
||||
[0.0, 0.0, 0.0, 1.0, 0.0],
|
||||
[0.0, 0.0, 0.0, 0.0, 1.0],
|
||||
],
|
||||
dtype="float32",
|
||||
)
|
||||
assertEqWithPresence(
|
||||
df["next_state_features"], next_state_features_presence, next_state_features
|
||||
)
|
||||
|
||||
assertEq(df["next_action"], np.array([1, 2, 3, 4]))
|
||||
assertEq(df["time_diff"], np.array([1, 3, 1, 1]))
|
||||
assertEq(df["step"], np.array([1, 1, 1, 1]))
|
||||
assertEq(
|
||||
df["possible_actions_mask"],
|
||||
np.array([[1, 1, 0, 0], [0, 1, 1, 0], [0, 0, 1, 1], [0, 0, 0, 1]]),
|
||||
)
|
||||
assertEq(
|
||||
df["possible_next_actions_mask"],
|
||||
np.array([[0, 1, 1, 0], [0, 0, 1, 1], [0, 0, 0, 1], [0, 0, 0, 0]]),
|
||||
)
|
||||
|
||||
|
||||
def verify_multi_step_except_rewards(df):
|
||||
assertEq(df["sequence_number"], pandas.Series([1, 2, 3, 4]))
|
||||
|
||||
state_features_presence = np.array(
|
||||
[
|
||||
[True, False, False, False, False],
|
||||
[False, True, False, False, False],
|
||||
[False, False, True, False, False],
|
||||
[False, False, False, True, False],
|
||||
],
|
||||
dtype="bool",
|
||||
)
|
||||
assertEq(df["state_features_presence"], state_features_presence)
|
||||
state_features = np.array(
|
||||
[
|
||||
[1.0, 0.0, 0.0, 0.0, 0.0],
|
||||
[0.0, 1.0, 0.0, 0.0, 0.0],
|
||||
[0.0, 0.0, 1.0, 0.0, 0.0],
|
||||
[0.0, 0.0, 0.0, 1.0, 0.0],
|
||||
],
|
||||
dtype="float32",
|
||||
)
|
||||
assertEqWithPresence(df["state_features"], state_features_presence, state_features)
|
||||
|
||||
assertEq(df["action"], np.array([0, 1, 2, 3]))
|
||||
assertEq(df["action_probability"], np.array([0.3, 0.4, 0.5, 0.6], dtype="float32"))
|
||||
assertEq(df["not_terminal"], np.array([1, 1, 0, 0], dtype="bool"))
|
||||
|
||||
next_state_features_presence = np.array(
|
||||
[
|
||||
[False, False, True, False, False],
|
||||
[False, False, False, True, False],
|
||||
[False, False, False, False, True],
|
||||
[False, False, False, False, True],
|
||||
],
|
||||
dtype="bool",
|
||||
)
|
||||
assertEq(df["next_state_features_presence"], next_state_features_presence)
|
||||
next_state_features = np.array(
|
||||
[
|
||||
[0.0, 0.0, 1.0, 0.0, 0.0],
|
||||
[0.0, 0.0, 0.0, 1.0, 0.0],
|
||||
[0.0, 0.0, 0.0, 0.0, 1.0],
|
||||
[0.0, 0.0, 0.0, 0.0, 1.0],
|
||||
],
|
||||
dtype="float32",
|
||||
)
|
||||
assertEqWithPresence(
|
||||
df["next_state_features"], next_state_features_presence, next_state_features
|
||||
)
|
||||
|
||||
assertEq(df["next_action"], np.array([2, 3, 4, 4]))
|
||||
assertEq(df["time_diff"], np.array([1, 1, 1, 1]))
|
||||
assertEq(df["step"], np.array([2, 2, 2, 1]))
|
||||
assertEq(
|
||||
df["possible_actions_mask"],
|
||||
np.array([[1, 1, 0, 0], [0, 1, 1, 0], [0, 0, 1, 1], [0, 0, 0, 1]]),
|
||||
)
|
||||
assertEq(
|
||||
df["possible_next_actions_mask"],
|
||||
np.array([[0, 0, 1, 1], [0, 0, 0, 1], [0, 0, 0, 0], [0, 0, 0, 0]]),
|
||||
)
|
||||
|
||||
|
||||
def rand_string(length=10):
|
||||
import string
|
||||
import random
|
||||
|
||||
"""Generate a random string of fixed length """
|
||||
letters = string.ascii_lowercase
|
||||
return "".join(random.choice(letters) for _ in range(length))
|
||||
|
||||
|
||||
class TestQueryData(ReagentSQLTestBase):
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
logging.getLogger(__name__).setLevel(logging.INFO)
|
||||
self.table_name = rand_string()
|
||||
self.temp_parquet = tempfile.TemporaryDirectory()
|
||||
self.parquet_url = f"file://{abspath(self.temp_parquet.name)}"
|
||||
logger.info(f"Table name is {self.table_name}")
|
||||
|
||||
def tearDown(self):
|
||||
self.temp_parquet.cleanup()
|
||||
super().tearDown()
|
||||
|
||||
def generate_data(self, multi_step=False):
|
||||
generate_data_discrete(
|
||||
self.sqlCtx, multi_step=multi_step, table_name=self.table_name
|
||||
)
|
||||
|
||||
def _read_data(self, custom_reward_expr=None, gamma=None, multi_step=None):
|
||||
query_data(
|
||||
table_spec=TableSpec(table_name=self.table_name),
|
||||
output_spec=Dataset(parquet_url=self.parquet_url),
|
||||
state_keys=[0, 1, 4, 5, 6],
|
||||
actions=["L", "R", "U", "D"],
|
||||
metrics_keys=["reward"],
|
||||
custom_reward_expr=custom_reward_expr,
|
||||
multi_step=multi_step,
|
||||
gamma=gamma,
|
||||
)
|
||||
df = self.sqlCtx.read.parquet(self.parquet_url)
|
||||
df = df.orderBy(asc("sequence_number"))
|
||||
logger.info("Read parquet dataframe")
|
||||
df.show()
|
||||
return df
|
||||
|
||||
@pytest.mark.serial
|
||||
def test_query_data_single_step(self):
|
||||
self.generate_data()
|
||||
df = self._read_data()
|
||||
df = df.toPandas()
|
||||
verify_single_step_except_rewards(df)
|
||||
assertEq(df["reward"], np.array([0.0, 1.0, 4.0, 5.0], dtype="float32"))
|
||||
logger.info("single-step seems fine")
|
||||
|
||||
@pytest.mark.serial
|
||||
def test_query_data_single_step_custom_reward(self):
|
||||
self.generate_data()
|
||||
df = self._read_data(custom_reward_expr="POWER(reward, 3) + 10")
|
||||
df = df.toPandas()
|
||||
verify_single_step_except_rewards(df)
|
||||
assertEq(df["reward"], np.array([10.0, 11.0, 74.0, 135.0], dtype="float32"))
|
||||
logger.info("single-step custom reward seems fine")
|
||||
|
||||
@pytest.mark.serial
|
||||
def test_query_data_multi_step(self):
|
||||
gamma = 0.9
|
||||
self.generate_data(multi_step=True)
|
||||
df = self._read_data(multi_step=2, gamma=gamma)
|
||||
df = df.toPandas()
|
||||
verify_multi_step_except_rewards(df)
|
||||
assertAllClose(
|
||||
df["reward"],
|
||||
np.array(
|
||||
[gamma * 1, 1 * 1.0 + gamma * 4, 1 * 4.0 + gamma * 5, 1 * 5.0],
|
||||
dtype="float32",
|
||||
),
|
||||
)
|
||||
logger.info("multi-step seems fine.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,5 +1,244 @@
|
||||
#!/usr/bin/env python3
|
||||
import logging
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from pyspark.sql.functions import col, crc32, udf
|
||||
from pyspark.sql.types import (
|
||||
ArrayType,
|
||||
BooleanType,
|
||||
FloatType,
|
||||
LongType,
|
||||
MapType,
|
||||
StructField,
|
||||
StructType,
|
||||
)
|
||||
from reagent.workflow.spark_utils import get_spark_session
|
||||
from reagent.workflow.types import Dataset, TableSpec
|
||||
|
||||
|
||||
def query_data(*args, **kwargs):
|
||||
raise NotImplementedError
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# for normalizing crc32 output
|
||||
MAX_UINT32 = 4294967295
|
||||
|
||||
|
||||
def calc_custom_reward(sqlCtx, df, custom_reward_expr: str):
|
||||
temp_table_name = "_tmp_calc_reward_df"
|
||||
temp_reward_name = "_tmp_reward_col"
|
||||
df.createOrReplaceTempView(temp_table_name)
|
||||
df = sqlCtx.sql(
|
||||
f"SELECT *, CAST(COALESCE({custom_reward_expr}, 0) AS FLOAT)"
|
||||
f" as {temp_reward_name} FROM {temp_table_name}"
|
||||
)
|
||||
return df.drop("reward").withColumnRenamed(temp_reward_name, "reward")
|
||||
|
||||
|
||||
def calc_reward_multi_step(sqlCtx, df, multi_step: int, gamma: float):
|
||||
# computes r_0 + gamma * (r_1 + gamma * (r_2 + ... ))
|
||||
expr = f"AGGREGATE(REVERSE(reward), FLOAT(0), (s, x) -> FLOAT({gamma}) * s + x)"
|
||||
return calc_custom_reward(sqlCtx, df, expr)
|
||||
|
||||
|
||||
def perform_preprocessing(
|
||||
sqlCtx,
|
||||
table_spec: TableSpec,
|
||||
state_keys: List[int],
|
||||
actions: List[str],
|
||||
metrics_keys: List[str],
|
||||
custom_reward_expr: Optional[str] = None,
|
||||
sample_range: Optional[Tuple[float, float]] = None,
|
||||
multi_step: Optional[int] = None,
|
||||
gamma: Optional[float] = None,
|
||||
):
|
||||
""" Perform preprocessing of given dataframe df.
|
||||
Preprocessing steps include calculating the reward,
|
||||
performing sparse-to-dense for mapped columns like state_features
|
||||
and metrics, and subsampling based on sample_range.
|
||||
If multi_step is set (with gamma), then we assume multi_step RL setting.
|
||||
"""
|
||||
if sample_range:
|
||||
assert (
|
||||
0.0 <= sample_range[0]
|
||||
and sample_range[0] <= sample_range[1]
|
||||
and sample_range[1] <= 100.0
|
||||
), f"{sample_range} is invalid."
|
||||
|
||||
df = sqlCtx.sql(f"SELECT * FROM {table_spec.table_name}")
|
||||
|
||||
# after this, reward column should be set to be the reward now
|
||||
if custom_reward_expr is not None:
|
||||
df = calc_custom_reward(sqlCtx, df, custom_reward_expr)
|
||||
elif multi_step is not None:
|
||||
assert gamma is not None
|
||||
df = calc_reward_multi_step(sqlCtx, df, multi_step, gamma)
|
||||
# assume single step case reward is already a column
|
||||
|
||||
def get_step(next_col):
|
||||
""" get step count """
|
||||
if multi_step is not None:
|
||||
return min(len(next_col), multi_step)
|
||||
else:
|
||||
return 1
|
||||
|
||||
get_step_udf = udf(get_step, LongType())
|
||||
df = df.withColumn("step", get_step_udf("next_state_features"))
|
||||
|
||||
def make_next_udf(return_type):
|
||||
""" return udf to get next item, provided item type """
|
||||
|
||||
def get_next(next_col):
|
||||
""" generic function to get the next item """
|
||||
if multi_step is not None:
|
||||
step = min(len(next_col), multi_step)
|
||||
return next_col[step - 1]
|
||||
else:
|
||||
return next_col
|
||||
|
||||
return udf(get_next, return_type)
|
||||
|
||||
df = df.withColumn("time_diff", make_next_udf(LongType())("time_diff"))
|
||||
|
||||
def make_sparse2dense(df, col_name: str, possible_keys: List):
|
||||
""" Given a list of possible keys, convert sparse map to dense array.
|
||||
In our example, both value_type is assumed to be a float.
|
||||
"""
|
||||
output_type = StructType(
|
||||
[
|
||||
StructField("presence", ArrayType(BooleanType()), False),
|
||||
StructField("dense", ArrayType(FloatType()), False),
|
||||
]
|
||||
)
|
||||
|
||||
def sparse2dense(map_col):
|
||||
assert isinstance(
|
||||
map_col, dict
|
||||
), f"{map_col} has type {type(map_col)} and is not a dict."
|
||||
presence = []
|
||||
dense = []
|
||||
for key in possible_keys:
|
||||
val = map_col.get(key, None)
|
||||
if val is not None:
|
||||
presence.append(True)
|
||||
dense.append(float(val))
|
||||
else:
|
||||
presence.append(False)
|
||||
dense.append(0.0)
|
||||
return presence, dense
|
||||
|
||||
sparse2dense_udf = udf(sparse2dense, output_type)
|
||||
df = df.withColumn(col_name, sparse2dense_udf(col_name))
|
||||
df = df.withColumn(f"{col_name}_presence", col(f"{col_name}.presence"))
|
||||
df = df.withColumn(col_name, col(f"{col_name}.dense"))
|
||||
return df
|
||||
|
||||
df = make_sparse2dense(df, "state_features", state_keys)
|
||||
|
||||
next_map_udf = make_next_udf(MapType(LongType(), FloatType()))
|
||||
df = df.withColumn("next_state_features", next_map_udf("next_state_features"))
|
||||
df = make_sparse2dense(df, "next_state_features", state_keys)
|
||||
|
||||
df = df.withColumn("metrics", next_map_udf("metrics"))
|
||||
df = make_sparse2dense(df, "metrics", metrics_keys)
|
||||
|
||||
def where(arr: List[str]):
|
||||
""" locate the index of item in arr, len(arr) if not found. """
|
||||
|
||||
def find(item: str):
|
||||
for i, arr_item in enumerate(arr):
|
||||
if arr_item == item:
|
||||
return i
|
||||
return len(arr)
|
||||
|
||||
return find
|
||||
|
||||
where_udf = udf(where(actions), LongType())
|
||||
df = df.withColumn("action", where_udf("action"))
|
||||
df = df.withColumn(
|
||||
"next_action", where_udf(make_next_udf(LongType())("next_action"))
|
||||
)
|
||||
|
||||
def get_not_terminal(next_action):
|
||||
""" terminal state iff next_action is "" (i.e. onehot len(actions))"""
|
||||
return next_action < len(actions)
|
||||
|
||||
get_not_terminal_udf = udf(get_not_terminal, BooleanType())
|
||||
df = df.withColumn("not_terminal", get_not_terminal_udf("next_action"))
|
||||
|
||||
def onehot(arr: List[str]):
|
||||
""" one-hot encode elements of arr depending on their existence in target """
|
||||
|
||||
def encode(target: List[str]):
|
||||
result = [0] * len(arr)
|
||||
for i, arr_item in enumerate(arr):
|
||||
if arr_item in target:
|
||||
result[i] = 1
|
||||
return result
|
||||
|
||||
return encode
|
||||
|
||||
onehot_udf = udf(onehot(actions), ArrayType(LongType()))
|
||||
df = df.withColumn("possible_actions_mask", onehot_udf("possible_actions"))
|
||||
df = df.withColumn(
|
||||
"possible_next_actions_mask",
|
||||
onehot_udf(make_next_udf(ArrayType(LongType()))("possible_next_actions")),
|
||||
)
|
||||
|
||||
# assuming use_seq_num_diff_as_time_diff = False for now
|
||||
df = df.withColumn("sequence_number", col("sequence_number_ordinal"))
|
||||
|
||||
# crc32 is treated as a cryptographic hash with range [0, MAX_UINT32-1]
|
||||
# Note: we're assuming no collisions!
|
||||
df = df.withColumn("mdp_id", crc32(col("mdp_id")))
|
||||
if sample_range:
|
||||
lower_bound = sample_range[0] / 100.0 * MAX_UINT32
|
||||
upper_bound = sample_range[1] / 100.0 * MAX_UINT32
|
||||
df = df.filter((lower_bound <= col("mdp_id")) & (col("mdp_id") <= upper_bound))
|
||||
|
||||
# select all the relevant columns and perform type conversions
|
||||
return df.select(
|
||||
col("reward").cast(FloatType()),
|
||||
col("state_features").cast(ArrayType(FloatType())),
|
||||
col("state_features_presence").cast(ArrayType(BooleanType())),
|
||||
col("next_state_features").cast(ArrayType(FloatType())),
|
||||
col("next_state_features_presence").cast(ArrayType(BooleanType())),
|
||||
col("action").cast(LongType()),
|
||||
col("action_probability").cast(FloatType()),
|
||||
col("not_terminal").cast(BooleanType()),
|
||||
col("next_action").cast(LongType()),
|
||||
col("possible_actions_mask").cast(ArrayType(LongType())),
|
||||
col("possible_next_actions_mask").cast(ArrayType(LongType())),
|
||||
col("mdp_id").cast(LongType()),
|
||||
col("sequence_number").cast(LongType()),
|
||||
col("step").cast(LongType()),
|
||||
col("time_diff").cast(LongType()),
|
||||
col("metrics").cast(ArrayType(FloatType())),
|
||||
col("metrics_presence").cast(ArrayType(BooleanType())),
|
||||
)
|
||||
|
||||
|
||||
def query_data(
|
||||
table_spec: TableSpec,
|
||||
output_spec: Dataset,
|
||||
state_keys: List[int],
|
||||
actions: List[str],
|
||||
metrics_keys: List[str],
|
||||
custom_reward_expr: Optional[str] = None,
|
||||
sample_range: Optional[Tuple[float, float]] = None,
|
||||
multi_step: Optional[int] = None,
|
||||
gamma: Optional[float] = None,
|
||||
) -> None:
|
||||
sqlCtx = get_spark_session()
|
||||
# includes rewards preprocessing, sparse2dense
|
||||
preprocessed_df = perform_preprocessing(
|
||||
sqlCtx,
|
||||
table_spec=table_spec,
|
||||
state_keys=state_keys,
|
||||
actions=actions,
|
||||
metrics_keys=metrics_keys,
|
||||
custom_reward_expr=custom_reward_expr,
|
||||
sample_range=sample_range,
|
||||
multi_step=multi_step,
|
||||
gamma=gamma,
|
||||
)
|
||||
preprocessed_df.write.mode("overwrite").parquet(output_spec.parquet_url)
|
||||
return
|
||||
|
||||
@@ -5,23 +5,14 @@ from typing import Dict, List, Optional
|
||||
|
||||
import pyspark
|
||||
from pyspark.sql.functions import col, collect_list, explode, flatten
|
||||
from reagent.preprocessing import normalization
|
||||
from reagent.preprocessing.normalization import (
|
||||
NormalizationParameters,
|
||||
get_feature_norm_metadata,
|
||||
)
|
||||
from reagent.workflow.spark_utils import get_spark_session
|
||||
from reagent.workflow.types import PreprocessingOptions, TableSpec
|
||||
|
||||
|
||||
LOCAL_MASTER = "local[1]"
|
||||
|
||||
|
||||
def get_spark_session(master: str = LOCAL_MASTER):
|
||||
""" Get a spark session """
|
||||
spark = (
|
||||
pyspark.sql.SparkSession.builder.master(master)
|
||||
.enableHiveSupport()
|
||||
.getOrCreate()
|
||||
)
|
||||
return spark
|
||||
|
||||
|
||||
def normalization_helper(
|
||||
max_unique_enum_values: int,
|
||||
quantile_size: int,
|
||||
@@ -47,7 +38,7 @@ def normalization_helper(
|
||||
whitelist_features = set(whitelist_features or [])
|
||||
|
||||
def validate_whitelist_features(
|
||||
params: Dict[int, normalization.NormalizationParameters],
|
||||
params: Dict[int, NormalizationParameters],
|
||||
) -> None:
|
||||
if not whitelist_features:
|
||||
return
|
||||
@@ -61,12 +52,12 @@ def normalization_helper(
|
||||
)
|
||||
)
|
||||
|
||||
def process(rows: List) -> Dict[int, normalization.NormalizationParameters]:
|
||||
def process(rows: List) -> Dict[int, NormalizationParameters]:
|
||||
params = {}
|
||||
for row in rows:
|
||||
assert "feature_name" in row
|
||||
assert "feature_values" in row
|
||||
norm_metdata = normalization.get_feature_norm_metadata(
|
||||
norm_metdata = get_feature_norm_metadata(
|
||||
row["feature_name"], row["feature_values"], norm_params
|
||||
)
|
||||
if norm_metdata is not None and (
|
||||
@@ -86,16 +77,14 @@ def identify_normalization_parameters(
|
||||
column_name: str,
|
||||
preprocessing_options: PreprocessingOptions,
|
||||
seed: int,
|
||||
) -> Dict[int, normalization.NormalizationParameters]:
|
||||
) -> Dict[int, NormalizationParameters]:
|
||||
""" Get normalization parameters """
|
||||
|
||||
spark = get_spark_session()
|
||||
df = spark.sql(f"SELECT * FROM {table_spec.table_name}")
|
||||
sqlCtx = get_spark_session()
|
||||
df = sqlCtx.sql(f"SELECT * FROM {table_spec.table_name}")
|
||||
df = create_normalization_spec_spark(
|
||||
df, column_name, preprocessing_options.num_samples, seed
|
||||
)
|
||||
rows = df.collect()
|
||||
spark.stop()
|
||||
|
||||
normalization_processor = normalization_helper(
|
||||
max_unique_enum_values=preprocessing_options.max_unique_enum_values,
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
import pyspark
|
||||
|
||||
|
||||
LOCAL_MASTER = "local[1]"
|
||||
|
||||
|
||||
def get_spark_session(master: str = LOCAL_MASTER):
|
||||
spark = (
|
||||
pyspark.sql.SparkSession.builder.master(master)
|
||||
.enableHiveSupport()
|
||||
.getOrCreate()
|
||||
)
|
||||
return spark
|
||||
@@ -24,13 +24,13 @@ from reagent.workflow.tagged_union import TaggedUnion # noqa F401
|
||||
|
||||
|
||||
@dataclass
|
||||
class TableSpec(BaseDataClass):
|
||||
class TableSpec:
|
||||
table_name: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class Dataset:
|
||||
url: str
|
||||
parquet_url: str
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -8,6 +8,7 @@ pyspark==2.4.5
|
||||
pytest==5.3
|
||||
pytest-xdist==1.30.0
|
||||
ruamel.yaml==0.15.99
|
||||
spark-testing-base==0.10.0
|
||||
scipy==1.3.1
|
||||
tensorboard==1.14
|
||||
scikit-learn==0.20.0
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
[tool:pytest]
|
||||
addopts = --verbose -d --tx 2*popen
|
||||
addopts = --verbose -d
|
||||
python_files = reagent/test/*.py reagent/test/**/*.py
|
||||
|
||||
[metadata]
|
||||
|
||||
@@ -13,5 +13,7 @@ deps =
|
||||
pytest-xdist==1.30.0
|
||||
# Pinning due to https://github.com/pytest-dev/pytest/issues/6925
|
||||
pytest==5.3
|
||||
spark-testing-base==0.10.0
|
||||
commands =
|
||||
pytest --junitxml={envlogdir}/junit-{envname}.xml
|
||||
pytest --junitxml={envlogdir}/junit-{envname}.xml -n auto --tx 2*popen -m "not serial"
|
||||
pytest --junitxml={envlogdir}/junit-{envname}-serial.xml -n0 --tx 1*popen -m "serial"
|
||||
|
||||
Reference in New Issue
Block a user