From d08c729bad4d7cbd8a674b648b9cb4490cc2ce07 Mon Sep 17 00:00:00 2001 From: Ruiyang Xu Date: Mon, 17 Aug 2020 00:09:17 -0700 Subject: [PATCH] CompressSeq2RewardModel Summary: Compress Seq2Reward Model through supervised learning another model Reviewed By: czxttkl Differential Revision: D22826516 fbshipit-source-id: b20bd0582ef895aef228be827889511580e8c84e --- .../evaluation/compress_model_evaluator.py | 27 ++++ reagent/net_builder/value/fully_connected.py | 4 +- reagent/parameters.py | 1 + reagent/prediction/predictor_wrapper.py | 23 ++-- reagent/training/utils.py | 11 ++ .../world_model/compress_model_trainer.py | 117 ++++++++++++++++++ .../world_model/seq2reward_trainer.py | 24 ++-- .../model_based/seq2reward_model.py | 7 ++ 8 files changed, 185 insertions(+), 29 deletions(-) create mode 100644 reagent/evaluation/compress_model_evaluator.py create mode 100644 reagent/training/world_model/compress_model_trainer.py diff --git a/reagent/evaluation/compress_model_evaluator.py b/reagent/evaluation/compress_model_evaluator.py new file mode 100644 index 00000000..f163563b --- /dev/null +++ b/reagent/evaluation/compress_model_evaluator.py @@ -0,0 +1,27 @@ +#!/usr/bin/env python3 +# Copyright (c) Facebook, Inc. and its affiliates. All rights reserved. +import logging + +import torch +from reagent.training.world_model.compress_model_trainer import CompressModelTrainer +from reagent.types import MemoryNetworkInput + + +logger = logging.getLogger(__name__) + + +class CompressModelEvaluator: + def __init__(self, trainer: CompressModelTrainer) -> None: + self.trainer = trainer + self.compress_model_network = self.trainer.compress_model_network + + # pyre-fixme[56]: Decorator `torch.no_grad(...)` could not be called, because + # its type `no_grad` is not callable. + @torch.no_grad() + def evaluate(self, eval_tdp: MemoryNetworkInput): + prev_mode = self.compress_model_network.training + self.compress_model_network.eval() + loss = self.trainer.get_loss(eval_tdp) + detached_loss = loss.cpu().detach().item() + self.compress_model_network.train(prev_mode) + return detached_loss diff --git a/reagent/net_builder/value/fully_connected.py b/reagent/net_builder/value/fully_connected.py index a8c491e1..cdf4157c 100644 --- a/reagent/net_builder/value/fully_connected.py +++ b/reagent/net_builder/value/fully_connected.py @@ -26,13 +26,13 @@ class FullyConnected(ValueNetBuilder): ) def build_value_network( - self, state_normalization_data: NormalizationData + self, state_normalization_data: NormalizationData, output_dim: int = 1 ) -> torch.nn.Module: state_dim = get_num_output_features( state_normalization_data.dense_normalization_parameters ) return FullyConnectedNetwork( - [state_dim] + self.sizes + [1], + [state_dim] + self.sizes + [output_dim], self.activations + ["linear"], use_layer_norm=self.use_layer_norm, ) diff --git a/reagent/parameters.py b/reagent/parameters.py index eb3c6b60..635fd8b9 100644 --- a/reagent/parameters.py +++ b/reagent/parameters.py @@ -73,6 +73,7 @@ class Seq2RewardTrainerParameters(BaseDataClass): action_names: List[str] = field(default_factory=lambda: []) batch_size: int = 32 gamma: float = 0.9 + view_q_value: bool = False @dataclass(frozen=True) diff --git a/reagent/prediction/predictor_wrapper.py b/reagent/prediction/predictor_wrapper.py index 20be728e..ea0db9dc 100644 --- a/reagent/prediction/predictor_wrapper.py +++ b/reagent/prediction/predictor_wrapper.py @@ -6,7 +6,6 @@ from typing import Dict, List, Optional, Tuple import reagent.types as rlt import torch -import torch.nn.functional as F from reagent.models.base import ModelBase from reagent.models.seq2slate import Seq2SlateMode, Seq2SlateTransformerNet from reagent.models.seq2slate_reward import Seq2SlateRewardNetBase @@ -17,6 +16,7 @@ from reagent.preprocessing.sparse_preprocessor import ( make_sparse_preprocessor, ) from reagent.torch_utils import gather +from reagent.training.utils import gen_permutations from torch import nn @@ -417,17 +417,6 @@ class Seq2RewardWithPreprocessor(DiscreteDqnWithPreprocessor): super().__init__(model, state_preprocessor, rlt.ModelFeatureConfig()) self.seq_len = seq_len self.num_action = num_action - - def gen_permutations(seq_len: int, num_action: int) -> torch.Tensor: - """ - generate all seq_len permutations for a given action set - the return shape is (SEQ_LEN, PERM_NUM, ACTION_DIM) - """ - all_permut = torch.cartesian_prod(*[torch.arange(num_action)] * seq_len) - all_permut = F.one_hot(all_permut, num_action).transpose(0, 1) - - return all_permut.float() - self.all_permut = gen_permutations(seq_len, num_action) self.num_permut = self.all_permut.size(1) @@ -607,3 +596,13 @@ class MDNRNNWithPreprocessor(ModelBase): self.state_preprocessor.input_prototype(), torch.randn(1, 1, self.num_action, device=self.state_preprocessor.device), ) + + +class CompressModelWithPreprocessor(DiscreteDqnWithPreprocessor): + def forward(self, state: rlt.ServingFeatureData): + state_feature_data = serving_to_feature_data( + state, self.state_preprocessor, self.sparse_preprocessor + ) + # TODO: model is a fully connected network which only takes in Tensor now. + q_values = self.model(state_feature_data.float_features) + return q_values diff --git a/reagent/training/utils.py b/reagent/training/utils.py index 81705dbf..03384916 100644 --- a/reagent/training/utils.py +++ b/reagent/training/utils.py @@ -5,6 +5,7 @@ from typing import Union import numpy as np import torch +import torch.nn.functional as F EPS = np.finfo(float).eps.item() @@ -48,3 +49,13 @@ def discounted_returns(rewards: torch.Tensor, gamma: float = 0) -> torch.Tensor: R = r + gamma * R returns.insert(0, R) return torch.tensor(returns).float() + + +def gen_permutations(seq_len: int, num_action: int) -> torch.Tensor: + """ + generate all seq_len permutations for a given action set + the return shape is (SEQ_LEN, PERM_NUM, ACTION_DIM) + """ + all_permut = torch.cartesian_prod(*[torch.arange(num_action)] * seq_len) + all_permut = F.one_hot(all_permut, num_action).transpose(0, 1) + return all_permut.float() diff --git a/reagent/training/world_model/compress_model_trainer.py b/reagent/training/world_model/compress_model_trainer.py new file mode 100644 index 00000000..cf631c12 --- /dev/null +++ b/reagent/training/world_model/compress_model_trainer.py @@ -0,0 +1,117 @@ +#!/usr/bin/env python3 +# Copyright (c) Facebook, Inc. and its affiliates. All rights reserved. + +import logging + +import reagent.types as rlt +import torch +import torch.nn.functional as F +from reagent.models.fully_connected_network import FullyConnectedNetwork +from reagent.models.seq2reward_model import Seq2RewardNetwork +from reagent.parameters import Seq2RewardTrainerParameters +from reagent.training.loss_reporter import NoOpLossReporter +from reagent.training.trainer import Trainer +from reagent.training.utils import gen_permutations + + +logger = logging.getLogger(__name__) + + +class CompressModelTrainer(Trainer): + """ Trainer for Seq2Reward """ + + def __init__( + self, + compress_model_network: FullyConnectedNetwork, + seq2reward_network: Seq2RewardNetwork, + params: Seq2RewardTrainerParameters, + ): + self.compress_model_network = compress_model_network + self.seq2reward_network = seq2reward_network + self.params = params + self.optimizer = torch.optim.Adam( + self.compress_model_network.parameters(), lr=params.learning_rate + ) + self.minibatch_size = self.params.batch_size + self.loss_reporter = NoOpLossReporter() + + # PageHandler must use this to activate evaluator: + self.calc_cpe_in_training = True + + def train(self, training_batch: rlt.MemoryNetworkInput): + self.optimizer.zero_grad() + loss = self.get_loss(training_batch) + loss.backward() + self.optimizer.step() + detached_loss = loss.cpu().detach().item() + + return detached_loss + + def get_loss(self, training_batch: rlt.MemoryNetworkInput): + compress_model_output = self.compress_model_network( + training_batch.state.float_features[0] + ) + target = self.get_Q( + training_batch, + training_batch.batch_size(), + self.params.multi_steps, + len(self.params.action_names), + ) + assert ( + compress_model_output.size() == target.size() + ), f"{compress_model_output.size()}!={target.size()}" + mse = F.mse_loss(compress_model_output, target) + return mse + + def warm_start_components(self): + logger.info("No warm start components yet...") + components = [] + return components + + # pyre-fixme[56]: Decorator `torch.no_grad(...)` could not be called, because + # its type `no_grad` is not callable. + @torch.no_grad() + def get_Q( + self, + batch: rlt.MemoryNetworkInput, + batch_size: int, + seq_len: int, + num_action: int, + ) -> torch.Tensor: + try: + # pyre-fixme[16]: `Seq2RewardTrainer` has no attribute `all_permut`. + self.all_permut + except AttributeError: + self.all_permut = gen_permutations(seq_len, num_action) + # pyre-fixme[16]: `Seq2RewardTrainer` has no attribute `num_permut`. + self.num_permut = self.all_permut.size(1) + + preprocessed_state = ( + batch.state.float_features[0] + .unsqueeze(0) + .repeat_interleave(self.num_permut, dim=1) + ) + state_feature_vector = rlt.FeatureData(preprocessed_state) + + # expand action to match the expanded state sequence + action = self.all_permut.repeat(1, batch_size, 1) + # state_feature_vector: [1, BATCH_SIZE * NUM_PERMUT, STATE_DIM] + # action: [SEQ_LEN, BATCH_SIZE * NUM_PERMUT, ACTION_DIM] + # acc_reward: [BATCH_SIZE * NUM_PERMUT, 1] + reward = self.seq2reward_network( + state_feature_vector, rlt.FeatureData(action) + ).acc_reward.reshape(batch_size, num_action, self.num_permut // num_action) + + # The permuations are generated with lexical order + # the output has shape [num_perm, num_action,1] + # that means we can aggregate on the max reward + # then reshape it to (BATCH_SIZE, ACT_DIM) + max_reward = ( + # pyre-fixme[16]: `Tuple` has no attribute `values`. + torch.max(reward, 2) + .values.cpu() + .detach() + .reshape(batch_size, num_action) + ) + + return max_reward diff --git a/reagent/training/world_model/seq2reward_trainer.py b/reagent/training/world_model/seq2reward_trainer.py index e2ec5f8e..db5259b3 100644 --- a/reagent/training/world_model/seq2reward_trainer.py +++ b/reagent/training/world_model/seq2reward_trainer.py @@ -10,6 +10,7 @@ from reagent.models.seq2reward_model import Seq2RewardNetwork from reagent.parameters import Seq2RewardTrainerParameters from reagent.training.loss_reporter import NoOpLossReporter from reagent.training.trainer import Trainer +from reagent.training.utils import gen_permutations logger = logging.getLogger(__name__) @@ -32,7 +33,7 @@ class Seq2RewardTrainer(Trainer): # PageHandler must use this to activate evaluator: self.calc_cpe_in_training = True # Turning off Q value output during training: - self.view_q_value = False + self.view_q_value = params.view_q_value def train(self, training_batch: rlt.MemoryNetworkInput): self.optimizer.zero_grad() @@ -81,13 +82,14 @@ class Seq2RewardTrainer(Trainer): target_acc_reward = torch.sum(target_rewards * gamma_mask, 0).unsqueeze(1) # make sure the prediction and target tensors have the same size # the size should both be (BATCH_SIZE, 1) in this case. - assert predicted_acc_reward.size() == target_acc_reward.size() + assert ( + predicted_acc_reward.size() == target_acc_reward.size() + ), f"{predicted_acc_reward.size()}!={target_acc_reward.size()}" mse = F.mse_loss(predicted_acc_reward, target_acc_reward) return mse def warm_start_components(self): - logger.info("No warm start components yet...") - components = [] + components = ["seq2reward_network"] return components def get_Q( @@ -103,21 +105,13 @@ class Seq2RewardTrainer(Trainer): # pyre-fixme[16]: `Seq2RewardTrainer` has no attribute `all_permut`. self.all_permut except AttributeError: - - def gen_permutations(seq_len: int, num_action: int) -> torch.Tensor: - """ - generate all seq_len permutations for a given action set - the return shape is (SEQ_LEN, PERM_NUM, ACTION_DIM) - """ - all_permut = torch.cartesian_prod(*[torch.arange(num_action)] * seq_len) - all_permut = F.one_hot(all_permut, num_action).transpose(0, 1) - return all_permut.float() - self.all_permut = gen_permutations(seq_len, num_action) # pyre-fixme[16]: `Seq2RewardTrainer` has no attribute `num_permut`. self.num_permut = self.all_permut.size(1) - preprocessed_state = batch.state.float_features.repeat(1, self.num_permut, 1) + preprocessed_state = batch.state.float_features.repeat_interleave( + self.num_permut, dim=1 + ) state_feature_vector = rlt.FeatureData(preprocessed_state) # expand action to match the expanded state sequence diff --git a/reagent/workflow/model_managers/model_based/seq2reward_model.py b/reagent/workflow/model_managers/model_based/seq2reward_model.py index cf749828..b48e8a96 100644 --- a/reagent/workflow/model_managers/model_based/seq2reward_model.py +++ b/reagent/workflow/model_managers/model_based/seq2reward_model.py @@ -5,6 +5,7 @@ import logging import torch from reagent.core.dataclasses import dataclass, field from reagent.net_builder.unions import ValueNetBuilder__Union +from reagent.net_builder.value.fully_connected import FullyConnected from reagent.net_builder.value.seq2reward_rnn import Seq2RewardNetBuilder from reagent.parameters import Seq2RewardTrainerParameters, param_hash from reagent.training.world_model.seq2reward_trainer import Seq2RewardTrainer @@ -25,6 +26,12 @@ class Seq2RewardModel(WorldModelBase): ) ) + compress_net_builder: ValueNetBuilder__Union = field( + # pyre-fixme[28]: Unexpected keyword argument `FullyConnected`. + # pyre-fixme[28]: Unexpected keyword argument `FullyConnected`. + default_factory=lambda: ValueNetBuilder__Union(FullyConnected=FullyConnected()) + ) + trainer_param: Seq2RewardTrainerParameters = field( default_factory=Seq2RewardTrainerParameters )