Ruurd commited on
Commit
8deda8e
路
verified 路
1 Parent(s): 63a94ec

Deploy BYOD-Llama-3.1-8B full-precision demo

Browse files
Files changed (2) hide show
  1. app.py +1 -1
  2. src/diffusion_lm/inference.py +9 -6
app.py CHANGED
@@ -75,7 +75,7 @@ def generate(
75
  permanent_unmask=True,
76
  confidence_guided=True,
77
  proportional_unmask=False,
78
- early_stopping=False,
79
  # This delays retention of predicted endings; it does not alter their
80
  # sampling probability. Keep it optional so answers can end naturally.
81
  confidence_eos_eot_inf=bool(delay_eos_eot),
 
75
  permanent_unmask=True,
76
  confidence_guided=True,
77
  proportional_unmask=False,
78
+ early_stopping=True,
79
  # This delays retention of predicted endings; it does not alter their
80
  # sampling probability. Keep it optional so answers can end naturally.
81
  confidence_eos_eot_inf=bool(delay_eos_eot),
src/diffusion_lm/inference.py CHANGED
@@ -985,11 +985,14 @@ def denoise_stream(session: InferenceSession, question: str, system_prompt: str,
985
  show_eos_tokens=include_pre_remask_prediction,
986
  )
987
  last_confidence = float(confidence[block_start:block_end].mean().cpu())
988
- # Match the legacy application's criterion: compare complete sampled
989
- # answer token sequences before the next iteration's re-masking. This
990
- # includes EOS/padding tokens, so a changing invisible tail does not
991
- # count as convergence.
992
- last_predictions.append(tuple(ids[answer_start + block_start : answer_start + block_end]))
 
 
 
993
  if len(last_predictions) > 3:
994
  last_predictions.pop(0)
995
  stopped_early = early_stopping and len(last_predictions) == 3 and len(set(last_predictions)) == 1
@@ -1079,7 +1082,7 @@ def denoise_stream(session: InferenceSession, question: str, system_prompt: str,
1079
  if permanent_unmask:
1080
  status += f" 路 retained {len(retained)} tokens"
1081
  if stopped_early:
1082
- status += " 路 block stopped early (same prediction for 3 iterations)"
1083
  html = render_denoising_step(
1084
  ids,
1085
  confidence.tolist(),
 
985
  show_eos_tokens=include_pre_remask_prediction,
986
  )
987
  last_confidence = float(confidence[block_start:block_end].mean().cpu())
988
+ # Compare the visible sampled answer before the next iteration's
989
+ # re-masking. Tokens after the first EOS are not part of the answer and
990
+ # must not prevent convergence. Excluding EOS itself still preserves
991
+ # its position through the tuple length: moving EOS changes the prefix.
992
+ prediction = ids[answer_start:]
993
+ if session.tokenizer.eos_token_id in prediction:
994
+ prediction = prediction[:prediction.index(session.tokenizer.eos_token_id)]
995
+ last_predictions.append(tuple(prediction))
996
  if len(last_predictions) > 3:
997
  last_predictions.pop(0)
998
  stopped_early = early_stopping and len(last_predictions) == 3 and len(set(last_predictions)) == 1
 
1082
  if permanent_unmask:
1083
  status += f" 路 retained {len(retained)} tokens"
1084
  if stopped_early:
1085
+ status += " 路 block stopped early (same answer for 3 iterations)"
1086
  html = render_denoising_step(
1087
  ids,
1088
  confidence.tolist(),