CompressSeq2RewardModel

Summary: Compress Seq2Reward Model through supervised learning another model

Reviewed By: czxttkl

Differential Revision: D22826516

fbshipit-source-id: b20bd0582ef895aef228be827889511580e8c84e
This commit is contained in:
Ruiyang Xu
2020-08-17 00:10:33 -07:00
committed by Facebook GitHub Bot
parent 192172311c
commit d08c729bad
8 changed files with 185 additions and 29 deletions
@@ -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
+2 -2
View File
@@ -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,
)
+1
View File
@@ -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)
+11 -12
View File
@@ -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
+11
View File
@@ -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()
@@ -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
@@ -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
@@ -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
)