import torch from transformers import ( DiffusionGemmaForBlockDiffusion, DiffusionGemmaGenerationConfig, EntropyBoundSamplerConfig, PreTrainedTokenizerFast, ) MODEL_PATH = "." MODEL_SUBFOLDER = "hf" PROMPT = "Once upon" def main() -> None: device = torch.device("cuda" if torch.cuda.is_available() else "cpu") tokenizer = PreTrainedTokenizerFast.from_pretrained( MODEL_PATH, subfolder=MODEL_SUBFOLDER, ) model = DiffusionGemmaForBlockDiffusion.from_pretrained( MODEL_PATH, subfolder=MODEL_SUBFOLDER, dtype=torch.float32, ).to(device) model.eval() generation_config = DiffusionGemmaGenerationConfig( max_new_tokens=64, max_denoising_steps=64, sampler_config=EntropyBoundSamplerConfig(entropy_bound=1.0), t_min=0.4, t_max=0.8, stability_threshold=3, confidence_threshold=0.05, bos_token_id=tokenizer.bos_token_id, eos_token_id=tokenizer.eos_token_id, pad_token_id=tokenizer.pad_token_id, cache_implementation="dynamic", return_dict_in_generate=True, ) input_ids = torch.tensor( [ [tokenizer.bos_token_id] + tokenizer.encode(PROMPT, add_special_tokens=False) ], dtype=torch.long, device=device, ) with torch.no_grad(): output = model.generate( input_ids=input_ids, generation_config=generation_config, ) sequences = output.sequences if hasattr(output, "sequences") else output print(tokenizer.decode(sequences[0].tolist(), skip_special_tokens=True)) if __name__ == "__main__": main()