x8D-Omni-Diffusion / byte_diffusion.py
bapX's picture
Upload folder using huggingface_hub
ddf1085 verified
Raw
History Blame Contribute Delete
11.8 kB
# coding=utf-8
"""DiffusionGemma-style byte diffusion sampler (pure stdlib).
Implements the x8D mapping of DiffusionGemma's uniform-state masked byte
diffusion over the 264-id byte vocabulary (bytes 0-255 + specials 256-263):
- ``canvas_length=256`` mirrors DiffusionGemma's ``canvas_length``.
- ``diffusion_entropy_bound`` mirrors DiffusionGemma's
``diffusion_entropy_bound`` sampler: positions whose denoiser confidence
falls below the bound are regenerated (block-autoregressive canvas
commit).
- ``sample_canvas`` iterates mask -> renoise -> denoise and commits a
fixed-size canvas, exactly the contract the torch DreamModel denoiser
will implement.
Deterministic given ``seed``; pure Python standard library only (no torch,
no transformers, no tokenizer vocabulary).
"""
from __future__ import annotations
import random
from typing import List, Optional, Tuple, Union
from .models.dream.byte_tokenizer import (
BOS_TOKEN_ID,
EOS_TOKEN_ID,
MASK_TOKEN_ID,
PAD_TOKEN_ID,
SPECIAL_IDS,
ByteTokenizer,
)
#: Stable mixing constant (FNV-style offset basis).
_HASH_OFFSET: int = 0x9E3779B9
#: Stable mixing constant (FNV-style prime).
_HASH_PRIME: int = 0x27D4EB2F
#: ASCII printable range used for regenerated bytes so the decoded output is
#: always valid UTF-8 (single-byte sequences, no dangling continuation bytes).
_ASCII_BASE: int = 0x20
_ASCII_COUNT: int = 0x5F # 0x20..0x7E inclusive -> 95 printable chars
def stable_hash(*values: int) -> int:
"""Deterministic, hash-independent integer mixing across process runs.
Uses only arithmetic (no ``hash()`` of tuples/strings), so results are
identical regardless of ``PYTHONHASHSEED``. Stand-in for the future torch
denoiser's learned positional/scalar mixing.
Args:
values: integer values to mix (e.g. ``(pos, byte, step)``).
Returns:
A 32-bit unsigned pseudo-random integer.
"""
h: int = _HASH_OFFSET
for v in values:
h = ((h ^ (int(v) & 0xFFFFFFFF)) * _HASH_PRIME + (h >> 13)) & 0xFFFFFFFF
h ^= h >> 16
h = (h * 0x85EBCA6B) & 0xFFFFFFFF
h ^= h >> 13
h = (h * 0xC2B2AE35) & 0xFFFFFFFF
h ^= h >> 16
return h
def _ascii_byte(position: int, step: int) -> int:
"""Printable-ASCII byte (0x20..0x7E) derived from a stable hash.
Guaranteed valid single-byte UTF-8, so any canvas made of these plus
original ASCII content decodes without a ``UnicodeDecodeError``.
Args:
position: canvas position index.
step: current diffusion step.
Returns:
A byte value in ``[0x20, 0x7E]``.
"""
return _ASCII_BASE + (stable_hash(position, step) % _ASCII_COUNT)
class ByteDiffusionSampler:
"""Pure-Python reference for the byte diffusion denoise loop.
Mirrors the masked-diffusion contract of DiffusionGemma's entropy-bound
sampler over the x8D 264-id byte space:
1. ``encode`` frames input as ``[BOS] <raw bytes> [EOS]``.
2. ``mask_canvas`` applies uniform-state masking (fraction of
positions -> MASK, id 256).
3. ``renoise_to_random_bytes`` refills masked positions with random
byte ids (seeded uniform-state re-noise).
4. ``denoise_step`` scores every position with a deterministic
pseudo-logit (stable-hash stand-in for the torch denoiser).
5. ``sample_canvas`` iterates mask -> renoise -> denoise, commits a
``canvas_length`` canvas, and returns the decoded bytes.
Deterministic given ``seed``.
"""
def __init__(
self,
vocab_size: int = 264,
canvas_length: int = 256,
diffusion_entropy_bound: float = 0.1,
seed: int = 0,
) -> None:
"""Initialize the sampler.
Args:
vocab_size: total vocabulary size (default 264 = 256 bytes + 8
specials). Must be at least 264 to cover specials 256-263.
canvas_length: fixed canvas length (default 256, DiffusionGemma
parity).
diffusion_entropy_bound: entropy-bound threshold; positions whose
denoiser confidence is below this are regenerated (default
0.1, DiffusionGemma parity).
seed: RNG seed for all masking/re-noise randomness.
"""
if vocab_size < 264:
raise ValueError(f"vocab_size must be >= 264 to cover specials, got {vocab_size}")
if canvas_length < 1:
raise ValueError(f"canvas_length must be >= 1, got {canvas_length}")
if not 0.0 <= diffusion_entropy_bound <= 1.0:
raise ValueError(
f"diffusion_entropy_bound must be in [0, 1], got {diffusion_entropy_bound}"
)
self.vocab_size: int = vocab_size
self.canvas_length: int = canvas_length
self.diffusion_entropy_bound: float = diffusion_entropy_bound
self.seed: int = seed
self._rng = random.Random(seed)
self._step: int = 0
self._decoder = ByteTokenizer()
def encode(self, text_or_bytes: Union[str, bytes, bytearray, List[int]]) -> List[int]:
"""Encode text/bytes into byte ids wrapped with BOS(258)/EOS(259).
Content ids are always the raw 0-255 byte states; there is no
vocabulary lookup or encoding step.
Args:
text_or_bytes: UTF-8 text, raw bytes, or an iterable of byte
ints.
Returns:
``[BOS, *byte_ids, EOS]``.
"""
if isinstance(text_or_bytes, str):
raw = text_or_bytes.encode("utf-8")
elif isinstance(text_or_bytes, bytearray):
raw = bytes(text_or_bytes)
elif isinstance(text_or_bytes, bytes):
raw = text_or_bytes
else:
raw = bytes(int(b) & 0xFF for b in text_or_bytes)
return [BOS_TOKEN_ID, *raw, EOS_TOKEN_ID]
def decode(
self,
ids: Union[List[int], bytes],
skip_special_tokens: bool = True,
as_bytes: bool = True,
) -> Union[bytes, str]:
"""Decode byte ids back into bytes (or UTF-8 text).
Args:
ids: iterable of byte ids (0-263).
skip_special_tokens: drop ids >= 256 when True (default).
as_bytes: return ``bytes`` when True (default), else UTF-8 str.
Returns:
The reconstructed byte string, or its UTF-8 text.
"""
return self._decoder.decode(ids, skip_special_tokens=skip_special_tokens, as_bytes=as_bytes)
def mask_canvas(
self, ids: List[int], mask_ratio: float = 0.7
) -> Tuple[List[int], List[int]]:
"""Uniform-state masking: mask a fraction of positions.
A ``round(len(ids) * mask_ratio)`` subset of positions is chosen
uniformly at random (seeded) and replaced with MASK(256).
Args:
ids: byte-id sequence to mask.
mask_ratio: fraction of positions to mask, in ``[0, 1]``.
Returns:
``(masked, truth)``: the masked ids and the untouched original.
"""
masked = list(ids)
truth = list(ids)
n = len(masked)
n_mask = int(round(n * mask_ratio))
n_mask = max(0, min(n_mask, n))
if n_mask:
positions = self._rng.sample(range(n), n_mask)
for p in positions:
masked[p] = MASK_TOKEN_ID
return masked, truth
def renoise_to_random_bytes(self, masked_ids: List[int]) -> List[int]:
"""Uniform-state re-noise: refill every MASK position with a random byte.
Args:
masked_ids: byte-id sequence possibly containing MASK(256).
Returns:
The sequence with every MASK position replaced by a seeded
uniform random byte id in 0-255.
"""
out = list(masked_ids)
for i, tok in enumerate(out):
if tok == MASK_TOKEN_ID:
out[i] = self._rng.randrange(256)
return out
def denoise_step(self, canvas: List[int]) -> List[float]:
"""Deterministic pseudo-logits per canvas position.
Stand-in for the future torch denoiser's output logits. Masked
positions score in ``[0, 0.5]`` (low confidence), committed content
positions score in ``(0.5, 1]``, both via ``stable_hash`` of
``(position, token, step)``. Position ``i`` of the returned list is
the pseudo-logit for canvas position ``i``.
Args:
canvas: byte-id canvas to score.
Returns:
One float pseudo-logit per position, each in ``[0, 1]``.
"""
logits: List[float] = []
for i, tok in enumerate(canvas):
h = stable_hash(i, tok if tok != MASK_TOKEN_ID else -1, self._step) & 0xFF
if tok == MASK_TOKEN_ID:
logits.append((h / 255.0) * 0.5)
else:
logits.append(0.5 + (h / 255.0) * 0.5)
return logits
def sample_canvas(
self,
prompt_bytes: Union[str, bytes, bytearray, List[int]],
steps: int = 8,
mask_ratio: float = 0.7,
entropy_bound: Optional[float] = None,
) -> bytes:
"""Canvas-based block sampling over the byte diffusion loop.
Iterates ``mask -> renoise -> denoise`` for ``steps`` rounds over a
fixed ``canvas_length`` canvas. Each round:
1. mask ``mask_ratio`` of the canvas (uniform state).
2. re-noise masked positions to random byte ids.
3. score positions with ``denoise_step``.
4. commit: confident positions (score >= entropy_bound) recover the
pre-mask byte when it is ASCII-safe; low-confidence positions are
regenerated from a stable hash into printable ASCII.
A final pass forces every remaining non-special byte into printable
ASCII so the decoded output is always valid UTF-8.
Args:
prompt_bytes: prompt text/bytes to seed the canvas.
steps: number of diffusion rounds.
mask_ratio: fraction of positions masked per round.
entropy_bound: threshold override; defaults to
``self.diffusion_entropy_bound``.
Returns:
Decoded content bytes (BOS/EOS/PAD/MASK and other specials
stripped).
"""
bound = self.diffusion_entropy_bound if entropy_bound is None else entropy_bound
if not 0.0 <= bound <= 1.0:
raise ValueError(f"entropy_bound must be in [0, 1], got {bound}")
# Reset the RNG so every call with this seed reproduces the exact
# same canvas sequence, regardless of prior calls on this instance.
self._rng = random.Random(self.seed)
self._step = 0
ids = self.encode(prompt_bytes)
if len(ids) > self.canvas_length:
keep = self.canvas_length - 2
ids = [BOS_TOKEN_ID, *ids[1 : 1 + keep], EOS_TOKEN_ID]
canvas = ids + [PAD_TOKEN_ID] * (self.canvas_length - len(ids))
for step in range(steps):
self._step = step
masked, truth = self.mask_canvas(canvas, mask_ratio=mask_ratio)
canvas = self.renoise_to_random_bytes(masked)
confidences = self.denoise_step(canvas)
for i in range(self.canvas_length):
if confidences[i] < bound:
canvas[i] = _ascii_byte(i, step)
else:
pre = truth[i]
if pre < 0x80:
canvas[i] = pre
for i in range(self.canvas_length):
b = canvas[i]
if b < 0x80 or b in SPECIAL_IDS:
continue
canvas[i] = _ascii_byte(i, self._step)
return bytes(self.decode(canvas, skip_special_tokens=True))