Add post_pos as a model target via travel_dir decomposition
step_length already encodes |post_pos - pre_pos| by definition, so a raw post_pos target would duplicate that magnitude and could drift inconsistent with step_length during sampling. Instead add travel_dir, a unit vector (local frame) giving only the direction of pre_pos->post_pos; post_pos is reconstructed at inference as pre_pos + step_length * travel_dir, keeping the two self-consistent. Target grows from 6D to 9D. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
+3
-3
@@ -10,7 +10,7 @@ def _small_model():
|
||||
|
||||
|
||||
def _batch(B=8):
|
||||
x1 = torch.randn(B, 6)
|
||||
x1 = torch.randn(B, 9)
|
||||
cond_cont = torch.randn(B, 9)
|
||||
cond_cat = torch.zeros(B, 2, dtype=torch.long)
|
||||
return x1, cond_cont, cond_cat
|
||||
@@ -40,7 +40,7 @@ def test_sample_flow_shape():
|
||||
cond_cont = torch.randn(B, 9)
|
||||
cond_cat = torch.zeros(B, 2, dtype=torch.long)
|
||||
out = sample_flow(_small_model(), cond_cont, cond_cat, steps=5)
|
||||
assert out.shape == (B, 6)
|
||||
assert out.shape == (B, 9)
|
||||
|
||||
|
||||
def test_ddpm_loss_nonneg():
|
||||
@@ -56,4 +56,4 @@ def test_sample_ddim_shape():
|
||||
cond_cont = torch.randn(B, 9)
|
||||
cond_cat = torch.zeros(B, 2, dtype=torch.long)
|
||||
out = sample_ddim(_small_model(), cond_cont, cond_cat, schedule, steps=5)
|
||||
assert out.shape == (B, 6)
|
||||
assert out.shape == (B, 9)
|
||||
|
||||
Reference in New Issue
Block a user