Skip to content

Commit 3ded40c

Browse files
David Vengerovfacebook-github-bot
authored andcommitted
Allow for publishing of reward network in discrete CRR
Summary: Allow for publishing of reward network in discrete_crr.py Differential Revision: D32711991 fbshipit-source-id: 7959d630ddce1e3c57688d9b449ffbeef7b9d8be
1 parent 4ab19c5 commit 3ded40c

2 files changed

Lines changed: 27 additions & 1 deletion

File tree

reagent/model_managers/discrete/discrete_crr.py

Lines changed: 25 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,12 @@
1515
param_hash,
1616
)
1717
from reagent.evaluation.evaluator import get_metrics_to_score
18+
from reagent.fb.prediction.cfeval.bandit_reward_network_predictor import (
19+
FbBanditRewardNetPredictorWrapper,
20+
)
21+
from reagent.fb.prediction.cfeval.bandit_reward_network_predictor import (
22+
FbBanditRewardNetPredictorWrapper,
23+
)
1824
from reagent.gym.policies.policy import Policy
1925
from reagent.gym.policies.predictor_policies import create_predictor_policy_from_model
2026
from reagent.model_managers.discrete_dqn_base import DiscreteDQNBase
@@ -196,7 +202,7 @@ def get_reporter(self):
196202
# in utils.py
197203

198204
def serving_module_names(self):
199-
module_names = ["default_model", "dqn", "actor_dqn"]
205+
module_names = ["default_model", "dqn", "actor_dqn", "reward"]
200206
if len(self.action_names) == 2:
201207
module_names.append("binary_difference_scorer")
202208
return module_names
@@ -219,6 +225,7 @@ def build_serving_modules(
219225
"dqn": self._build_dqn_module(
220226
trainer_module.q1_network, normalization_data_map
221227
),
228+
"reward": self.build_reward_module(trainer_module, normalization_data_map),
222229
"actor_dqn": self._build_dqn_module(
223230
ActorDQN(trainer_module.actor_network), normalization_data_map
224231
),
@@ -286,6 +293,23 @@ def build_actor_module(
286293
action_feature_ids=list(range(len(self.action_names))),
287294
)
288295

296+
def build_reward_module(
297+
self,
298+
trainer_module: DiscreteCRRTrainer,
299+
normalization_data_map: Dict[str, NormalizationData],
300+
) -> torch.nn.Module:
301+
"""
302+
Returns a TorchScript predictor module
303+
"""
304+
net_builder = self.cpe_net_builder.value
305+
return net_builder.build_serving_module(
306+
trainer_module.reward_network,
307+
normalization_data_map[NormalizationKey.STATE],
308+
action_names=self.action_names,
309+
state_feature_config=self.state_feature_config,
310+
predictor_wrapper_type=FbBanditRewardNetPredictorWrapper,
311+
)
312+
289313

290314
class ActorDQN(ModelBase):
291315
def __init__(self, actor):

reagent/training/discrete_crr_trainer.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -135,6 +135,8 @@ def __init__(
135135
# pyre-fixme[16]: Optional type has no attribute `__getitem__`.
136136
self.reward_boosts[0, i] = rl.reward_boost[k]
137137

138+
# The function below adds reward_network as a member object to DQNTrainerBaseLightning,
139+
# from which DiscreteCRRTrainer is derived.
138140
self._initialize_cpe(
139141
reward_network,
140142
q_network_cpe,

0 commit comments

Comments
 (0)