MCPcopy Create free account
hub / github.com/MotrixLab/ADHMR / generalized_steps

Function generalized_steps

ADHMR/lib/utils/diff_utils.py:66–97  ·  view source on GitHub ↗

x_0: joint gaussian noise X_T x_1: twist gaussian noise X_T

(x_0, x_1, seq, model, b, eta, ctx=None, gen_multi=False)

Source from the content-addressed store, hash-verified

64
65
66def generalized_steps(x_0, x_1, seq, model, b, eta, ctx=None, gen_multi=False):
67 '''
68 x_0: joint gaussian noise X_T
69 x_1: twist gaussian noise X_T
70 '''
71 with torch.no_grad():
72 n = x_0.size(0)
73 seq_next = [-1] + list(seq[:-1])
74 x0_preds = [[],[]]
75 xs = [[x_0],[x_1]]
76 for i, j in zip(reversed(seq), reversed(seq_next)):
77 t = (torch.ones(n) * i).to(x_0.device)
78 next_t = (torch.ones(n) * j).to(x_0.device)
79 at = compute_alpha(b, t.long())
80 at_next = compute_alpha(b, next_t.long())
81 xt_0 = xs[0][-1].to(x_0.device)
82 xt_1 = xs[1][-1].to(x_0.device)
83
84 et = model(xinj=xt_0,xint=xt_1,t=t.float(),ctx =ctx, gen_multi= gen_multi) # estimated noise for current timestep
85 x0_t_0 = (xt_0 - et[0] * (1 - at).sqrt()) / at.sqrt()
86 x0_t_1 = (xt_1 - et[1] * (1 - at).sqrt()) / at.sqrt()
87 x0_preds[0].append(x0_t_0) # estimated x_0 of current timestep
88 x0_preds[1].append(x0_t_1)
89 c1 = ( # signma_t
90 eta * ((1 - at / at_next) * (1 - at_next) / (1 - at)).sqrt()
91 )
92 c2 = ((1 - at_next) - c1 ** 2).sqrt()
93 xt_next_0 = at_next.sqrt() * x0_t_0 + c1 * torch.randn_like(x_0) + c2 * et[0]
94 xt_next_1 = at_next.sqrt() * x0_t_1 + c1 * torch.randn_like(x_1) + c2 * et[1]
95 xs[0].append(xt_next_0)
96 xs[1].append(xt_next_1)
97 return xs, x0_preds

Callers 3

sample_poseMethod · 0.85
sample_poseMethod · 0.85
sample_poseMethod · 0.85

Calls 2

compute_alphaFunction · 0.85
toMethod · 0.45

Tested by

no test coverage detected