66from ding .hpc_rl import hpc_wrapper
77
88ppo_data = namedtuple (
9- 'ppo_data' , ['logit_new' , 'logit_old' , 'action' , 'value_new' , 'value_old' , 'adv' , 'return_' , 'weight' ]
9+ 'ppo_data' ,
10+ ['logit_new' , 'logit_old' , 'action' , 'value_new' , 'value_old' , 'adv' , 'return_' , 'weight' , 'logit_pretrained' ]
1011)
1112ppo_data_continuous = namedtuple (
1213 'ppo_data_continuous' ,
1314 ['mu_sigma_new' , 'mu_sigma_old' , 'action' , 'value_new' , 'value_old' , 'adv' , 'return_' , 'weight' ]
1415)
15- ppo_policy_data = namedtuple ('ppo_policy_data' , ['logit_new' , 'logit_old' , 'action' , 'adv' , 'weight' ])
16+ ppo_policy_data = namedtuple (
17+ 'ppo_policy_data' , ['logit_new' , 'logit_old' , 'action' , 'adv' , 'weight' , 'logit_pretrained' ]
18+ )
1619ppo_policy_data_continuous = namedtuple (
1720 'ppo_policy_data_continuous' , ['mu_sigma_new' , 'mu_sigma_old' , 'action' , 'adv' , 'weight' ]
1821)
2225ppo_info = namedtuple ('ppo_info' , ['approx_kl' , 'clipfrac' , 'kl_div' ])
2326
2427
28+ def calculate_kl_div (logr : torch .Tensor , kl_type : str ) -> torch .Tensor :
29+ """
30+ Overview:
31+ Calculate different Monte-Carlo estimators for KL-divergence KL(q, p) = E_q[log(q/p)],
32+ where q is the current policy and p is the pretrained policy.
33+ The implementation is based on John Schulman's blog post "Approximating KL Divergence".
34+ Reference: http://joschu.net/blog/kl-approx.html
35+ Arguments:
36+ - logr (:obj:`torch.Tensor`): The log-ratio of probabilities, which should be log(q/p) = logp_new - logp_pretrained.
37+ - kl_type (:obj:`str`): The type of KL divergence estimator to use.
38+ - 'k1': The standard, unbiased but high-variance estimator: `E_q[log(q/p)]`.
39+ - 'k2': A biased, low-variance estimator from a second-order approximation: `E_q[1/2 * (log(p/q))^2]`.
40+ - 'k3': An unbiased, low-variance estimator: `E_q[(p/q - 1) - log(p/q)]`.
41+ Returns:
42+ - kl_div (:obj:`torch.Tensor`): The calculated KL divergence estimate.
43+ """
44+ if kl_type == 'k1' :
45+ return logr .mean ()
46+ elif kl_type == 'k2' :
47+ return (logr ** 2 / 2 ).mean ()
48+ elif kl_type == 'k3' :
49+ return (torch .exp (- logr ) - 1 + logr ).mean ()
50+ else :
51+ raise ValueError (f"Unknown kl_type: { kl_type } " )
52+
53+
2554def shape_fn_ppo (args , kwargs ):
2655 r"""
2756 Overview:
@@ -97,8 +126,8 @@ def ppo_error(
97126 assert dual_clip is None or dual_clip > 1.0 , "dual_clip value must be greater than 1.0, but get value: {}" .format (
98127 dual_clip
99128 )
100- logit_new , logit_old , action , value_new , value_old , adv , return_ , weight = data
101- policy_data = ppo_policy_data (logit_new , logit_old , action , adv , weight )
129+ logit_new , logit_old , action , value_new , value_old , adv , return_ , weight , logit_pretrained = data
130+ policy_data = ppo_policy_data (logit_new , logit_old , action , adv , weight , logit_pretrained )
102131 policy_output , policy_info = ppo_policy_error (policy_data , clip_ratio , dual_clip , kl_type = kl_type )
103132 value_data = ppo_value_data (value_new , value_old , return_ , weight )
104133 value_loss = ppo_value_error (value_data , clip_ratio , use_value_clip )
@@ -152,7 +181,7 @@ def ppo_policy_error(
152181 .. note::
153182 For the action mask often used in LLM/VLM, users can set the `weight` to the action mask.
154183 """
155- logit_new , logit_old , action , adv , weight = data
184+ logit_new , logit_old , action , adv , weight , logit_pretrained = data
156185 if weight is None :
157186 weight = torch .ones_like (adv )
158187 dist_new = torch .distributions .categorical .Categorical (logits = logit_new )
@@ -185,15 +214,13 @@ def ppo_policy_error(
185214 clipped = ratio .gt (1 + clip_ratio ) | ratio .lt (1 - clip_ratio )
186215 clipfrac = torch .as_tensor (clipped ).float ().mean ().item ()
187216
188- logr = logp_old - logp_new
189- if kl_type == 'k1' :
190- kl_div = logr .mean ()
191- elif kl_type == 'k2' :
192- kl_div = (logr ** 2 / 2 ).mean ()
193- elif kl_type == 'k3' :
194- kl_div = (torch .exp (- logr ) - 1 + logr ).mean ()
217+ if logit_pretrained is not None :
218+ dist_pretrained = torch .distributions .categorical .Categorical (logits = logit_pretrained )
219+ logp_pretrained = dist_pretrained .log_prob (action )
220+ logr = logp_new - logp_pretrained
221+ kl_div = calculate_kl_div (logr , kl_type )
195222 else :
196- raise ValueError ( f"Unknown kl_type: { kl_type } " )
223+ kl_div = 0
197224
198225 return ppo_policy_loss (policy_loss , entropy_loss ), ppo_info (approx_kl , clipfrac , kl_div )
199226
@@ -298,7 +325,7 @@ def ppo_error_continuous(
298325 assert dual_clip is None or dual_clip > 1.0 , "dual_clip value must be greater than 1.0, but get value: {}" .format (
299326 dual_clip
300327 )
301- mu_sigma_new , mu_sigma_old , action , value_new , value_old , adv , return_ , weight = data
328+ mu_sigma_new , mu_sigma_old , action , value_new , value_old , adv , return_ , weight , logit_pretrained = data
302329 if weight is None :
303330 weight = torch .ones_like (adv )
304331
@@ -331,15 +358,13 @@ def ppo_error_continuous(
331358 else :
332359 value_loss = 0.5 * ((return_ - value_new ).pow (2 ) * weight ).mean ()
333360
334- logr = logp_old - logp_new
335- if kl_type == 'k1' :
336- kl_div = logr .mean ()
337- elif kl_type == 'k2' :
338- kl_div = (logr ** 2 / 2 ).mean ()
339- elif kl_type == 'k3' :
340- kl_div = (torch .exp (- logr ) - 1 + logr ).mean ()
361+ if logit_pretrained is not None :
362+ dist_pretrained = Independent (Normal (logit_pretrained ['mu' ], logit_pretrained ['sigma' ]), 1 )
363+ logp_pretrained = dist_pretrained .log_prob (action )
364+ logr = logp_new - logp_pretrained
365+ kl_div = calculate_kl_div (logr , kl_type )
341366 else :
342- raise ValueError ( f"Unknown kl_type: { kl_type } " )
367+ kl_div = 0
343368
344369 return ppo_loss (policy_loss , value_loss , entropy_loss ), ppo_info (approx_kl , clipfrac , kl_div )
345370
@@ -384,7 +409,7 @@ def ppo_policy_error_continuous(
384409 assert dual_clip is None or dual_clip > 1.0 , "dual_clip value must be greater than 1.0, but get value: {}" .format (
385410 dual_clip
386411 )
387- mu_sigma_new , mu_sigma_old , action , adv , weight = data
412+ mu_sigma_new , mu_sigma_old , action , adv , weight , logit_pretrained = data
388413 if weight is None :
389414 weight = torch .ones_like (adv )
390415
@@ -409,14 +434,12 @@ def ppo_policy_error_continuous(
409434 clipped = ratio .gt (1 + clip_ratio ) | ratio .lt (1 - clip_ratio )
410435 clipfrac = torch .as_tensor (clipped ).float ().mean ().item ()
411436
412- logr = logp_old - logp_new
413- if kl_type == 'k1' :
414- kl_div = logr .mean ()
415- elif kl_type == 'k2' :
416- kl_div = (logr ** 2 / 2 ).mean ()
417- elif kl_type == 'k3' :
418- kl_div = (torch .exp (- logr ) - 1 + logr ).mean ()
437+ if logit_pretrained is not None :
438+ dist_pretrained = Independent (Normal (logit_pretrained ['mu' ], logit_pretrained ['sigma' ]), 1 )
439+ logp_pretrained = dist_pretrained .log_prob (action )
440+ logr = logp_new - logp_pretrained
441+ kl_div = calculate_kl_div (logr , kl_type )
419442 else :
420- raise ValueError ( f"Unknown kl_type: { kl_type } " )
443+ kl_div = 0
421444
422445 return ppo_policy_loss (policy_loss , entropy_loss ), ppo_info (approx_kl , clipfrac , kl_div )
0 commit comments