codewithdark/DiffusionLM
132
1import torch2import torch.nn as nn3from transformers import PretrainedConfig, PreTrainedModel4from diffusionLM.model.diffusionLM import LLaDAModel5 6class DiffusionConfig(PretrainedConfig):7 """Configuration class for Diffusion-LLM model."""8 model_type = "diffusionLM"9 10 def __init__(11 self,12 vocab_size: int = 50257,13 hidden_size: int = 768,14 num_hidden_layers: int = 12,15 num_attention_heads: int = 12,16 intermediate_size: int = 3072,17 hidden_dropout_prob: float = 0.1,18 attention_probs_dropout_prob: float = 0.1,19 max_position_embeddings: int = 1024,20 initializer_range: float = 0.02,21 layer_norm_eps: float = 1e-12,22 pad_token_id: int = 0,23 mask_token_id: int = 50256,24 eos_token_id: int = 50256,25 num_timesteps: int = 100,26 time_embed_dim: int = 128,27 **kwargs28 ):29 super().__init__(pad_token_id=pad_token_id, **kwargs)30 self.vocab_size = vocab_size31 self.hidden_size = hidden_size32 self.num_hidden_layers = num_hidden_layers33 self.num_attention_heads = num_attention_heads 34 self.intermediate_size = intermediate_size35 self.hidden_dropout_prob = hidden_dropout_prob36 self.attention_probs_dropout_prob = attention_probs_dropout_prob37 self.max_position_embeddings = max_position_embeddings38 self.initializer_range = initializer_range39 self.layer_norm_eps = layer_norm_eps40 self.mask_token_id = mask_token_id41 self.eos_token_id = eos_token_id42 self.num_timesteps = num_timesteps43 self.time_embed_dim = time_embed_dim44 45class DiffusionLLM(PreTrainedModel):46 """Main Diffusion-LLM model class"""47 config_class = DiffusionConfig48 base_model_prefix = "diffusionLM"49 50 def __init__(self, config: DiffusionConfig):51 super().__init__(config)52 self.model = LLaDAModel(config)53 self.init_weights()54 55 def forward(56 self,57 input_ids=None,58 attention_mask=None,59 timesteps=None,60 labels=None,61 return_dict=True,62 ):63 outputs = self.model(64 input_ids=input_ids,65 attention_mask=attention_mask,66 timesteps=timesteps,67 labels=labels,68 )69 70 return outputs71 72 def generate(73 self,74 prompt=None,75 max_length=100,76 num_inference_steps=50,77 temperature=1.0,78 strategy='random',79 top_p=0.9,80 top_k=50,81 num_beams=5,82 return_scores=False,83 use_streaming=False,84 callback_fn=None85 ):86 """Unified generation interface"""87 if use_streaming:88 return self.generate_stream(89 prompt=prompt,90 max_length=max_length,91 num_inference_steps=num_inference_steps,92 temperature=temperature,93 strategy=strategy,94 top_p=top_p,95 top_k=top_k,96 num_beams=num_beams,97 callback_fn=callback_fn98 )99 else:100 return self.model.generate(101 prompt=prompt,102 max_length=max_length,103 num_inference_steps=num_inference_steps,104 temperature=temperature,105 strategy=strategy,106 top_p=top_p,107 top_k=top_k,108 num_beams=num_beams,109 return_scores=return_scores110 )111 112 def generate_stream(self, **kwargs):113 """Streaming generation wrapper"""114 return self.model.generate_stream(**kwargs)115 116 def prepare_inputs_for_generation(self, input_ids, **kwargs):117 """Prepare inputs for generation compatibility"""118 return {119 "input_ids": input_ids,120 "attention_mask": kwargs.get("attention_mask", None),121 "timesteps": kwargs.get("timesteps", None),122 }123 124 @staticmethod125 def _reorder_cache(past, beam_idx):126 """Reorder cache for beam search compatibility"""127 return past128 