"""Local correction for dex 0.5.0 PositionOptimizer's inconsistent loss/gradient.""" import numpy as np from dex_retargeting.optimizer import PositionOptimizer class ConsistentPositionOptimizer(PositionOptimizer): def get_objective_function(self,target_pos,fixed_qpos,last_qpos): original=super().get_objective_function(target_pos,fixed_qpos,last_qpos) def objective(x,grad): # Upstream already includes this term's gradient, but omits its value. return original(x,grad)+self.norm_delta*float(np.sum((x-last_qpos)**2)) return objective