mirror of
https://github.com/facebookresearch/ReAgent.git
synced 2026-06-16 12:44:41 +00:00
Summary: Need more tests before landing the refactor diffs: D22702504 (https://github.com/facebookresearch/ReAgent/commit/1b470c489d19c33beab88b8ea2e79843d4d31f28), D23123762 (https://github.com/facebookresearch/ReAgent/commit/76829287265bc39f879f3bc1d946a1374c5e1141), D23124179 (https://github.com/facebookresearch/ReAgent/commit/b28f84aa013be00194508f52498160592cb37e9d), D23219012 (https://github.com/facebookresearch/ReAgent/commit/e404c5772ea4118105c2eb136ca96ad5ca8e01db) Back out to a version based on D23155753. Check our team diff history: https://fburl.com/diffs/ppsgazgj Reviewed By: kittipatv Differential Revision: D23270626 fbshipit-source-id: 14653066bb3924a987a54650a51241895b321c8e
57 lines
1.4 KiB
Python
57 lines
1.4 KiB
Python
#!/usr/bin/env python3
|
|
# Copyright (c) Facebook, Inc. and its affiliates. All rights reserved.
|
|
|
|
import dataclasses
|
|
import logging
|
|
import unittest
|
|
from typing import Any
|
|
|
|
import torch
|
|
import torch.nn as nn
|
|
from reagent import types as rlt
|
|
from reagent.models.base import ModelBase
|
|
from reagent.test.models.test_utils import check_save_load
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
@dataclasses.dataclass
|
|
class ModelOutput:
|
|
# These should be torch.Tensor but the type checking failed when I used it
|
|
sum: Any
|
|
mul: Any
|
|
plus_one: Any
|
|
linear: Any
|
|
|
|
|
|
class Model(ModelBase):
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.linear = nn.Linear(4, 1)
|
|
|
|
def input_prototype(self):
|
|
return (
|
|
rlt.FeatureData(torch.randn([1, 4])),
|
|
rlt.FeatureData(torch.randn([1, 4])),
|
|
)
|
|
|
|
def forward(self, state, action):
|
|
state = state.float_features
|
|
action = action.float_features
|
|
|
|
return ModelOutput(
|
|
state + action, state * action, state + 1, self.linear(state)
|
|
)
|
|
|
|
|
|
class TestBase(unittest.TestCase):
|
|
def test_get_predictor_export_meta_and_workspace(self):
|
|
model = Model()
|
|
|
|
# 2 params + 1 const
|
|
expected_num_params, expected_num_inputs, expected_num_outputs = 3, 2, 4
|
|
check_save_load(
|
|
self, model, expected_num_params, expected_num_inputs, expected_num_outputs
|
|
)
|