Team Ai
Modelpublic

codewithdark/DiffusionLM

sourceHugging Facemitupdated 1y agoView on Hugging Face
1likes32downloads
modeling_DiffusionLLM.py128 linesDownload Raw Back to root
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