Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 15 additions & 8 deletions hackable_diffusion/lib/sampling/diffusion_early_stopping.py
Original file line number Diff line number Diff line change
Expand Up @@ -122,6 +122,11 @@ def should_stop(
previous_step: DiffusionStep,
) -> Bool['B']: # pyrefly: ignore[not-a-type]
del step, previous_step
xt = current_step.xt
if len(xt.shape) != 3:
raise ValueError(
f'xt must have shape (batch_size, seq_len, 1) but got {xt.shape}'
)
aux = current_step.aux
logits = aux[self.logits_key]
log_probs = jax.nn.log_softmax(logits)
Expand Down Expand Up @@ -164,15 +169,17 @@ def should_stop(
) -> Bool['B']: # pyrefly: ignore[not-a-type]
del step
prev_tokens = previous_step.xt
assert len(prev_tokens.shape) == 3, (
'prev_tokens must have shape (batch_size, seq_len, 1) but got'
f' {prev_tokens.shape}'
)
if len(prev_tokens.shape) != 3:
raise ValueError(
'prev_tokens must have shape (batch_size, seq_len, 1) but got'
f' {prev_tokens.shape}'
)
batch_size, seq_len = prev_tokens.shape[:2]
assert prev_tokens.shape == (batch_size, seq_len, 1), (
'prev_tokens must have shape (batch_size, seq_len, 1) but got'
' {prev_tokens.shape}'
)
if prev_tokens.shape != (batch_size, seq_len, 1):
raise ValueError(
'prev_tokens must have shape (batch_size, seq_len, 1) but got'
f' {prev_tokens.shape}'
)
logits = current_step.aux[self.logits_key]
most_likely_tokens = jnp.argmax(logits, axis=-1)
prev_tokens = jnp.reshape(prev_tokens, (batch_size, seq_len))
Expand Down
100 changes: 75 additions & 25 deletions hackable_diffusion/lib/sampling/diffusion_early_stopping_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,9 +96,9 @@ def _make_step_with_logits(
logits_key: str = 'logits',
) -> hd_api.DiffusionStep:
"""Helper: creates a DiffusionStep with logits in aux."""
batch_size = logits.shape[0]
batch_size, seq_len = logits.shape[:2]
return _make_diffusion_step(
xt=jnp.ones((batch_size, 4)),
xt=jnp.ones((batch_size, seq_len, 1)),
aux={logits_key: logits},
)

Expand Down Expand Up @@ -248,21 +248,41 @@ def test_ignores_step_and_previous(self):
)
chex.assert_trees_all_equal(result_a, result_b)

def test_raises_value_error_for_invalid_xt_shape(self):
"""Should raise ValueError if xt length of shape is not 3."""
fn = diffusion_early_stopping.DiffusionEntropyEarlyStopFn()
step = _make_diffusion_step(
xt=jnp.ones((2, 4)), aux={'logits': jnp.zeros((2, 4, 3))}
)
with self.assertRaisesRegex(
ValueError, r'xt must have shape \(batch_size, seq_len, 1\) but got'
):
fn.should_stop(
step=jnp.int32(0),
current_step=step,
previous_step=step,
)


class DiffusionTokenStabilityEarlyStopFnTest(parameterized.TestCase):
"""Tests for DiffusionTokenStabilityEarlyStopFn."""

def setUp(self):
super().setUp()
prev_tokens = jnp.array([[0, 1, 2], [3, 2, 1]])
self.prev_tokens = jnp.reshape(prev_tokens, (2, 3, 1))

def test_stable_tokens_stop(self):
"""When argmax(logits) matches previous_step.xt, should stop."""
prev_tokens = jnp.array([[0, 1, 2], [3, 2, 1]])
prev_tokens = jnp.reshape(prev_tokens, (2, 3, 1))
logits = jnp.full((2, 3, 4), -100.0)
for b in range(2):
for l in range(3):
logits = logits.at[b, l, prev_tokens[b, l]].set(100.0)
logits = logits.at[b, l, self.prev_tokens[b, l]].set(100.0)

previous_step = _make_diffusion_step(xt=prev_tokens)
current_step = _make_diffusion_step(xt=prev_tokens, aux={'logits': logits})
previous_step = _make_diffusion_step(xt=self.prev_tokens)
current_step = _make_diffusion_step(
xt=self.prev_tokens, aux={'logits': logits}
)

fn = diffusion_early_stopping.DiffusionTokenStabilityEarlyStopFn()
result = fn.should_stop(
Expand All @@ -275,12 +295,12 @@ def test_stable_tokens_stop(self):

def test_unstable_tokens_continue(self):
"""When argmax(logits) differs from previous_step.xt, should not stop."""
prev_tokens = jnp.array([[0, 1, 2], [3, 2, 1]])
prev_tokens = jnp.reshape(prev_tokens, (2, 3, 1))
logits = jnp.full((2, 3, 4), -100.0).at[:, :, 3].set(100.0)

previous_step = _make_diffusion_step(xt=prev_tokens)
current_step = _make_diffusion_step(xt=prev_tokens, aux={'logits': logits})
previous_step = _make_diffusion_step(xt=self.prev_tokens)
current_step = _make_diffusion_step(
xt=self.prev_tokens, aux={'logits': logits}
)

fn = diffusion_early_stopping.DiffusionTokenStabilityEarlyStopFn()
result = fn.should_stop(
Expand All @@ -293,15 +313,15 @@ def test_unstable_tokens_continue(self):

def test_per_batch_mixed_stability(self):
"""Batch element 0 is stable, batch element 1 is unstable."""
prev_tokens = jnp.array([[0, 1, 2], [3, 2, 1]])
prev_tokens = jnp.reshape(prev_tokens, (2, 3, 1))
logits = jnp.full((2, 3, 4), -100.0)
for l in range(3):
logits = logits.at[0, l, prev_tokens[0, l]].set(100.0)
logits = logits.at[0, l, self.prev_tokens[0, l]].set(100.0)
logits = logits.at[1, :, 0].set(100.0)

previous_step = _make_diffusion_step(xt=prev_tokens)
current_step = _make_diffusion_step(xt=prev_tokens, aux={'logits': logits})
previous_step = _make_diffusion_step(xt=self.prev_tokens)
current_step = _make_diffusion_step(
xt=self.prev_tokens, aux={'logits': logits}
)

fn = diffusion_early_stopping.DiffusionTokenStabilityEarlyStopFn()
result = fn.should_stop(
Expand All @@ -315,15 +335,14 @@ def test_per_batch_mixed_stability(self):

def test_custom_logits_key(self):
"""Should read logits from custom logits_key in aux."""
prev_tokens = jnp.array([[0, 1, 2]])
prev_tokens = jnp.reshape(prev_tokens, (1, 3, 1))
logits = jnp.full((1, 3, 4), -100.0)
for l in range(3):
logits = logits.at[0, l, prev_tokens[0, l]].set(100.0)
logits = jnp.full((2, 3, 4), -100.0)
for b in range(2):
for l in range(3):
logits = logits.at[b, l, self.prev_tokens[b, l]].set(100.0)

previous_step = _make_diffusion_step(xt=prev_tokens)
previous_step = _make_diffusion_step(xt=self.prev_tokens)
current_step = _make_diffusion_step(
xt=prev_tokens, aux={'my_logits': logits}
xt=self.prev_tokens, aux={'my_logits': logits}
)

fn = diffusion_early_stopping.DiffusionTokenStabilityEarlyStopFn(
Expand All @@ -334,7 +353,38 @@ def test_custom_logits_key(self):
current_step=current_step,
previous_step=previous_step,
)
self.assertTrue(result[0])
self.assertTrue(jnp.all(result))

def test_raises_value_error_for_invalid_prev_tokens_shape(self):
"""Should raise ValueError if prev_tokens shape is not (batch_size, seq_len, 1)."""
fn = diffusion_early_stopping.DiffusionTokenStabilityEarlyStopFn()

# Case 1: 2D prev_tokens shape (2, 4)
prev_step_2d = _make_diffusion_step(xt=jnp.ones((2, 4)))
current_step = _make_diffusion_step(
xt=jnp.ones((2, 4, 1)), aux={'logits': jnp.zeros((2, 4, 3))}
)
with self.assertRaisesRegex(
ValueError,
r'prev_tokens must have shape \(batch_size, seq_len, 1\) but got',
):
fn.should_stop(
step=jnp.int32(0),
current_step=current_step,
previous_step=prev_step_2d,
)

# Case 2: 3D prev_tokens shape with trailing dim != 1 (2, 4, 2)
prev_step_bad_dim = _make_diffusion_step(xt=jnp.ones((2, 4, 2)))
with self.assertRaisesRegex(
ValueError,
r'prev_tokens must have shape \(batch_size, seq_len, 1\) but got',
):
fn.should_stop(
step=jnp.int32(0),
current_step=current_step,
previous_step=prev_step_bad_dim,
)


class DiffusionChainedEarlyStopFnTest(parameterized.TestCase):
Expand Down Expand Up @@ -381,7 +431,7 @@ def test_any_stopper_false_continues(self):
)

logits = jnp.zeros((2, 3, 4))
step = _make_diffusion_step(xt=jnp.ones((2, 4)), aux={'logits': logits})
step = _make_diffusion_step(xt=jnp.ones((2, 3, 1)), aux={'logits': logits})

result = chained.should_stop(
step=jnp.int32(0), current_step=step, previous_step=step
Expand Down
Loading