MCPcopy Create free account
hub / github.com/OpenDriveLab/DriveAdapter / _get_features

Method _get_features

roach/models/ppo_policy.py:102–109  ·  view source on GitHub ↗

:param birdview: th.Tensor (num_envs, frame_stack*channel, height, width) :param state: th.Tensor (num_envs, state_dim)

(self, birdview: th.Tensor, state: th.Tensor)

Source from the content-addressed store, hash-verified

100 self.optimizer = self.optimizer_class(self.parameters(), **self.optimizer_kwargs)
101
102 def _get_features(self, birdview: th.Tensor, state: th.Tensor) -> th.Tensor:
103 """
104 :param birdview: th.Tensor (num_envs, frame_stack*channel, height, width)
105 :param state: th.Tensor (num_envs, state_dim)
106 """
107 birdview = birdview.float() / 255.0
108 features, cnn_feature = self.features_extractor(birdview, state)
109 return features, cnn_feature
110
111 def _get_action_dist_from_features(self, features: th.Tensor):
112 latent_pi = self.policy_head(features)

Callers 5

evaluate_actionsMethod · 0.95
evaluate_valuesMethod · 0.95
forwardMethod · 0.95
forward_valueMethod · 0.95
forward_policyMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected