mirror of
https://github.com/facebookresearch/ReAgent.git
synced 2026-06-16 12:44:41 +00:00
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:
committed by
Facebook GitHub Bot
parent
192172311c
commit
d08c729bad
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user