Skip to content

Commit 1854e58

Browse files
committed
fix(nyz): fix ppo logit pretrained compatibility bugs
1 parent 8d29a32 commit 1854e58

3 files changed

Lines changed: 16 additions & 15 deletions

File tree

ding/policy/ppo.py

Lines changed: 13 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -328,15 +328,15 @@ def _forward_learn(self, data: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
328328
# discrete part (discrete policy loss and entropy loss)
329329
ppo_discrete_batch = ppo_policy_data(
330330
output['logit']['action_type'], batch['logit']['action_type'], batch['action']['action_type'],
331-
adv, batch['weight']
331+
adv, batch['weight'], logit_pretrained
332332
)
333333
ppo_discrete_loss, ppo_discrete_info = ppo_policy_error(
334334
ppo_discrete_batch, self._clip_ratio, kl_type=self._kl_type
335335
)
336336
# continuous part (continuous policy loss and entropy loss, value loss)
337337
ppo_continuous_batch = ppo_data(
338338
output['logit']['action_args'], batch['logit']['action_args'], batch['action']['action_args'],
339-
output['value'], batch['value'], adv, batch['return'], batch['weight']
339+
output['value'], batch['value'], adv, batch['return'], batch['weight'], None
340340
)
341341
ppo_continuous_loss, ppo_continuous_info = ppo_error_continuous(
342342
ppo_continuous_batch, self._clip_ratio, kl_type=self._kl_type
@@ -794,7 +794,7 @@ def _forward_learn(self, data: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
794794
output = self._learn_model.forward(batch['obs'])
795795

796796
ppo_batch = ppo_policy_data(
797-
output['logit'], batch['logit'], batch['action'], batch['return'], batch['weight']
797+
output['logit'], batch['logit'], batch['action'], batch['return'], batch['weight'], None
798798
)
799799
if self._action_space == 'continuous':
800800
ppo_loss, ppo_info = ppo_policy_error_continuous(ppo_batch, self._clip_ratio)
@@ -1235,26 +1235,26 @@ def _forward_learn(self, data: List[Dict[str, Any]]) -> Dict[str, Any]:
12351235
if self._action_space == 'continuous':
12361236
ppodata = ppo_data(
12371237
output['logit'], data['logit'], data['action'], output['value'], data['value'], adv, data['return'],
1238-
data['weight']
1238+
data['weight'], None
12391239
)
12401240
ppo_loss, ppo_info = ppo_error_continuous(ppodata, self._clip_ratio)
12411241
elif self._action_space == 'discrete':
12421242
ppodata = ppo_data(
12431243
output['logit'], data['logit'], data['action'], output['value'], data['value'], adv, data['return'],
1244-
data['weight']
1244+
data['weight'], None
12451245
)
12461246
ppo_loss, ppo_info = ppo_error(ppodata, self._clip_ratio)
12471247
elif self._action_space == 'hybrid':
12481248
# discrete part (discrete policy loss and entropy loss)
12491249
ppo_discrete_batch = ppo_policy_data(
12501250
output['logit']['action_type'], data['logit']['action_type'], data['action']['action_type'], adv,
1251-
data['weight']
1251+
data['weight'], None
12521252
)
12531253
ppo_discrete_loss, ppo_discrete_info = ppo_policy_error(ppo_discrete_batch, self._clip_ratio)
12541254
# continuous part (continuous policy loss and entropy loss, value loss)
12551255
ppo_continuous_batch = ppo_data(
12561256
output['logit']['action_args'], data['logit']['action_args'], data['action']['action_args'],
1257-
output['value'], data['value'], adv, data['return'], data['weight']
1257+
output['value'], data['value'], adv, data['return'], data['weight'], None
12581258
)
12591259
ppo_continuous_loss, ppo_continuous_info = ppo_error_continuous(ppo_continuous_batch, self._clip_ratio)
12601260
# sum discrete and continuous loss
@@ -1279,22 +1279,22 @@ def _forward_learn(self, data: List[Dict[str, Any]]) -> Dict[str, Any]:
12791279

12801280
# Calculate ppo loss
12811281
if self._action_space == 'continuous':
1282-
ppodata = ppo_policy_data(output['logit'], data['logit'], data['action'], adv, data['weight'])
1282+
ppodata = ppo_policy_data(output['logit'], data['logit'], data['action'], adv, data['weight'], None)
12831283
ppo_policy_loss, ppo_info = ppo_policy_error_continuous(ppodata, self._clip_ratio)
12841284
elif self._action_space == 'discrete':
1285-
ppodata = ppo_policy_data(output['logit'], data['logit'], data['action'], adv, data['weight'])
1285+
ppodata = ppo_policy_data(output['logit'], data['logit'], data['action'], adv, data['weight'], None)
12861286
ppo_policy_loss, ppo_info = ppo_policy_error(ppodata, self._clip_ratio)
12871287
elif self._action_space == 'hybrid':
12881288
# discrete part (discrete policy loss and entropy loss)
12891289
ppo_discrete_data = ppo_policy_data(
12901290
output['logit']['action_type'], data['logit']['action_type'], data['action']['action_type'], adv,
1291-
data['weight']
1291+
data['weight'], None
12921292
)
12931293
ppo_discrete_loss, ppo_discrete_info = ppo_policy_error(ppo_discrete_data, self._clip_ratio)
12941294
# continuous part (continuous policy loss and entropy loss, value loss)
12951295
ppo_continuous_data = ppo_policy_data(
12961296
output['logit']['action_args'], data['logit']['action_args'], data['action']['action_args'], adv,
1297-
data['weight']
1297+
data['weight'], None
12981298
)
12991299
ppo_continuous_loss, ppo_continuous_info = ppo_policy_error_continuous(
13001300
ppo_continuous_data, self._clip_ratio
@@ -1804,13 +1804,13 @@ def _forward_learn(self, data: Dict[str, Any]) -> Dict[str, Any]:
18041804
if self._action_space == 'continuous':
18051805
ppo_batch = ppo_data(
18061806
output['logit'], batch['logit'], batch['action'], output['value'], batch['value'], adv,
1807-
batch['return'], batch['weight']
1807+
batch['return'], batch['weight'], None
18081808
)
18091809
ppo_loss, ppo_info = ppo_error_continuous(ppo_batch, self._clip_ratio)
18101810
elif self._action_space == 'discrete':
18111811
ppo_batch = ppo_data(
18121812
output['logit'], batch['logit'], batch['action'], output['value'], batch['value'], adv,
1813-
batch['return'], batch['weight']
1813+
batch['return'], batch['weight'], None
18141814
)
18151815
ppo_loss, ppo_info = ppo_error(ppo_batch, self._clip_ratio)
18161816

ding/rl_utils/ppo.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -225,7 +225,7 @@ def ppo_policy_error(
225225
log_ratio = logp_new - logp_pretrained
226226
kl_div = calculate_kl_div(log_ratio, kl_type)
227227
else:
228-
kl_div = 0
228+
kl_div = torch.tensor(0., dtype=policy_loss.dtype, device=policy_loss.device)
229229

230230
return ppo_policy_loss(policy_loss, entropy_loss, kl_div), ppo_info(approx_kl, clipfrac)
231231

@@ -369,7 +369,7 @@ def ppo_error_continuous(
369369
log_ratio = logp_new - logp_pretrained
370370
kl_div = calculate_kl_div(log_ratio, kl_type)
371371
else:
372-
kl_div = 0
372+
kl_div = torch.tensor(0., dtype=policy_loss.dtype, device=policy_loss.device)
373373

374374
return ppo_loss(policy_loss, value_loss, entropy_loss, kl_div), ppo_info(approx_kl, clipfrac)
375375

ding/worker/learner/learner_hook.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -287,6 +287,7 @@ def should_reduce(key):
287287
# The "noreduce_" prefix is used in the unizero_multitask ddp pipeline
288288
# to indicate data that should not be reduced.
289289
return not key.startswith("noreduce_")
290+
290291
cuda_device = torch.cuda.current_device()
291292

292293
if isinstance(data, dict):

0 commit comments

Comments
 (0)