Team Ai
Apppublic

Zwounds/Boolean_Search_Query_Model

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
UnslothAlignPropTrainer.py638 linesDownload Raw Back to unsloth_compiled_cache
1"""22025.3.1532025.3.1744.50.0.dev050.15.26__UNSLOTH_VERSIONING__7"""8from torch import Tensor9import torch10import torch.nn as nn11from torch.nn import functional as F12from trl.trainer.alignprop_trainer import (Accelerator, AlignPropConfig, AlignPropTrainer, Any, Callable, DDPOStableDiffusionPipeline, Optional, ProjectConfiguration, PyTorchModelHubMixin, Union, defaultdict, generate_model_card, get_comet_experiment_url, is_wandb_available, logger, os, set_seed, textwrap, torch, warn)13 14 15import os16from typing import *17from dataclasses import dataclass, field18from packaging.version import Version19import torch20import numpy as np21from contextlib import nullcontext22from torch.nn import functional as F23from transformers import DataCollatorForSeq2Seq, DataCollatorForLanguageModeling24 25torch_compile_options = {26    "epilogue_fusion"   : True,27    "max_autotune"      : False,28    "shape_padding"     : True,29    "trace.enabled"     : False,30    "triton.cudagraphs" : False,31}32 33@torch.compile(dynamic = True, fullgraph = True, options = torch_compile_options,)34def selective_log_softmax(logits, index):35    logits = logits.to(torch.float32)36    selected_logits = torch.gather(logits, dim = -1, index = index.unsqueeze(-1)).squeeze(-1)37    # loop to reduce peak mem consumption38    # logsumexp_values = torch.stack([torch.logsumexp(lg, dim=-1) for lg in logits])39    logsumexp_values = torch.logsumexp(logits, dim = -1)40    per_token_logps = selected_logits - logsumexp_values  # log_softmax(x_i) = x_i - logsumexp(x)41    return per_token_logps42@dataclass43class UnslothAlignPropConfig(AlignPropConfig):44    """45    46    Configuration class for the [`AlignPropTrainer`].47 48    Using [`~transformers.HfArgumentParser`] we can turn this class into49    [argparse](https://docs.python.org/3/library/argparse#module-argparse) arguments that can be specified on the50    command line.51 52    Parameters:53        exp_name (`str`, *optional*, defaults to `os.path.basename(sys.argv[0])[: -len(".py")]`):54            Name of this experiment (defaults to the file name without the extension).55        run_name (`str`, *optional*, defaults to `""`):56            Name of this run.57        seed (`int`, *optional*, defaults to `0`):58            Random seed for reproducibility.59        log_with (`str` or `None`, *optional*, defaults to `None`):60            Log with either `"wandb"` or `"tensorboard"`. Check61            [tracking](https://huggingface.co/docs/accelerate/usage_guides/tracking) for more details.62        log_image_freq (`int`, *optional*, defaults to `1`):63            Frequency for logging images.64        tracker_kwargs (`dict[str, Any]`, *optional*, defaults to `{}`):65            Keyword arguments for the tracker (e.g., `wandb_project`).66        accelerator_kwargs (`dict[str, Any]`, *optional*, defaults to `{}`):67            Keyword arguments for the accelerator.68        project_kwargs (`dict[str, Any]`, *optional*, defaults to `{}`):69            Keyword arguments for the accelerator project config (e.g., `logging_dir`).70        tracker_project_name (`str`, *optional*, defaults to `"trl"`):71            Name of project to use for tracking.72        logdir (`str`, *optional*, defaults to `"logs"`):73            Top-level logging directory for checkpoint saving.74        num_epochs (`int`, *optional*, defaults to `100`):75            Number of epochs to train.76        save_freq (`int`, *optional*, defaults to `1`):77            Number of epochs between saving model checkpoints.78        num_checkpoint_limit (`int`, *optional*, defaults to `5`):79            Number of checkpoints to keep before overwriting old ones.80        mixed_precision (`str`, *optional*, defaults to `"fp16"`):81            Mixed precision training.82        allow_tf32 (`bool`, *optional*, defaults to `True`):83            Allow `tf32` on Ampere GPUs.84        resume_from (`str`, *optional*, defaults to `""`):85            Path to resume training from a checkpoint.86        sample_num_steps (`int`, *optional*, defaults to `50`):87            Number of sampler inference steps.88        sample_eta (`float`, *optional*, defaults to `1.0`):89            Eta parameter for the DDIM sampler.90        sample_guidance_scale (`float`, *optional*, defaults to `5.0`):91            Classifier-free guidance weight.92        train_batch_size (`int`, *optional*, defaults to `1`):93            Batch size for training.94        train_use_8bit_adam (`bool`, *optional*, defaults to `False`):95            Whether to use the 8bit Adam optimizer from `bitsandbytes`.96        train_learning_rate (`float`, *optional*, defaults to `1e-3`):97            Learning rate.98        train_adam_beta1 (`float`, *optional*, defaults to `0.9`):99            Beta1 for Adam optimizer.100        train_adam_beta2 (`float`, *optional*, defaults to `0.999`):101            Beta2 for Adam optimizer.102        train_adam_weight_decay (`float`, *optional*, defaults to `1e-4`):103            Weight decay for Adam optimizer.104        train_adam_epsilon (`float`, *optional*, defaults to `1e-8`):105            Epsilon value for Adam optimizer.106        train_gradient_accumulation_steps (`int`, *optional*, defaults to `1`):107            Number of gradient accumulation steps.108        train_max_grad_norm (`float`, *optional*, defaults to `1.0`):109            Maximum gradient norm for gradient clipping.110        negative_prompts (`str` or `None`, *optional*, defaults to `None`):111            Comma-separated list of prompts to use as negative examples.112        truncated_backprop_rand (`bool`, *optional*, defaults to `True`):113            If `True`, randomized truncation to different diffusion timesteps is used.114        truncated_backprop_timestep (`int`, *optional*, defaults to `49`):115            Absolute timestep to which the gradients are backpropagated. Used only if `truncated_backprop_rand=False`.116        truncated_rand_backprop_minmax (`tuple[int, int]`, *optional*, defaults to `(0, 50)`):117            Range of diffusion timesteps for randomized truncated backpropagation.118        push_to_hub (`bool`, *optional*, defaults to `False`):119            Whether to push the final model to the Hub.120    121    """122    vllm_sampling_params: Optional[Any] = field(123        default = None,124        metadata = {'help': 'vLLM SamplingParams'},125    )126    unsloth_num_chunks : Optional[int] = field(127        default = -1,128        metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},129    )130    def __init__(131        self,132        exp_name = 'demo',133        run_name = '',134        seed = 3407,135        log_with = None,136        log_image_freq = 1,137        tracker_project_name = 'trl',138        logdir = 'logs',139        num_epochs = 100,140        save_freq = 1,141        num_checkpoint_limit = 5,142        mixed_precision = 'fp16',143        allow_tf32 = True,144        resume_from = '',145        sample_num_steps = 50,146        sample_eta = 1.0,147        sample_guidance_scale = 5.0,148        train_batch_size = 1,149        train_use_8bit_adam = False,150        train_learning_rate = 5e-05,151        train_adam_beta1 = 0.9,152        train_adam_beta2 = 0.999,153        train_adam_weight_decay = 0.01,154        train_adam_epsilon = 1e-08,155        train_gradient_accumulation_steps = 2,156        train_max_grad_norm = 1.0,157        negative_prompts = None,158        truncated_backprop_rand = True,159        truncated_backprop_timestep = 49,160        push_to_hub = False,161        vllm_sampling_params = None,162        unsloth_num_chunks = -1,163        **kwargs,164    ):165        166        super().__init__(167            exp_name = exp_name,168            run_name = run_name,169            seed = seed,170            log_with = log_with,171            log_image_freq = log_image_freq,172            tracker_project_name = tracker_project_name,173            logdir = logdir,174            num_epochs = num_epochs,175            save_freq = save_freq,176            num_checkpoint_limit = num_checkpoint_limit,177            mixed_precision = mixed_precision,178            allow_tf32 = allow_tf32,179            resume_from = resume_from,180            sample_num_steps = sample_num_steps,181            sample_eta = sample_eta,182            sample_guidance_scale = sample_guidance_scale,183            train_batch_size = train_batch_size,184            train_use_8bit_adam = train_use_8bit_adam,185            train_learning_rate = train_learning_rate,186            train_adam_beta1 = train_adam_beta1,187            train_adam_beta2 = train_adam_beta2,188            train_adam_weight_decay = train_adam_weight_decay,189            train_adam_epsilon = train_adam_epsilon,190            train_gradient_accumulation_steps = train_gradient_accumulation_steps,191            train_max_grad_norm = train_max_grad_norm,192            negative_prompts = negative_prompts,193            truncated_backprop_rand = truncated_backprop_rand,194            truncated_backprop_timestep = truncated_backprop_timestep,195            push_to_hub = push_to_hub,**kwargs)196        self.vllm_sampling_params = vllm_sampling_params197        self.unsloth_num_chunks = unsloth_num_chunks198pass199 200class _UnslothAlignPropTrainer(PyTorchModelHubMixin):201    """"""202 203    _tag_names = ["trl", "alignprop"]204 205    def __init__(206        self,207        config: AlignPropConfig,208        reward_function: Callable[[torch.Tensor, tuple[str], tuple[Any]], torch.Tensor],209        prompt_function: Callable[[], tuple[str, Any]],210        sd_pipeline: DDPOStableDiffusionPipeline,211        image_samples_hook: Optional[Callable[[Any, Any, Any], Any]] = None,212    ):213        if image_samples_hook is None:214            warn("No image_samples_hook provided; no images will be logged")215 216        self.prompt_fn = prompt_function217        self.reward_fn = reward_function218        self.config = config219        self.image_samples_callback = image_samples_hook220 221        accelerator_project_config = ProjectConfiguration(**self.config.project_kwargs)222 223        if self.config.resume_from:224            self.config.resume_from = os.path.normpath(os.path.expanduser(self.config.resume_from))225            if "checkpoint_" not in os.path.basename(self.config.resume_from):226                # get the most recent checkpoint in this directory227                checkpoints = list(228                    filter(229                        lambda x: "checkpoint_" in x,230                        os.listdir(self.config.resume_from),231                    )232                )233                if len(checkpoints) == 0:234                    raise ValueError(f"No checkpoints found in {self.config.resume_from}")235                checkpoint_numbers = sorted([int(x.split("_")[-1]) for x in checkpoints])236                self.config.resume_from = os.path.join(237                    self.config.resume_from,238                    f"checkpoint_{checkpoint_numbers[-1]}",239                )240 241                accelerator_project_config.iteration = checkpoint_numbers[-1] + 1242 243        self.accelerator = Accelerator(244            log_with=self.config.log_with,245            mixed_precision=self.config.mixed_precision,246            project_config=accelerator_project_config,247            # we always accumulate gradients across timesteps; we want config.train.gradient_accumulation_steps to be the248            # number of *samples* we accumulate across, so we need to multiply by the number of training timesteps to get249            # the total number of optimizer steps to accumulate across.250            gradient_accumulation_steps=self.config.train_gradient_accumulation_steps,251            **self.config.accelerator_kwargs,252        )253 254        is_using_tensorboard = config.log_with is not None and config.log_with == "tensorboard"255 256        if self.accelerator.is_main_process:257            self.accelerator.init_trackers(258                self.config.tracker_project_name,259                config=dict(alignprop_trainer_config=config.to_dict())260                if not is_using_tensorboard261                else config.to_dict(),262                init_kwargs=self.config.tracker_kwargs,263            )264 265        logger.info(f"\n{config}")266 267        set_seed(self.config.seed, device_specific=True)268 269        self.sd_pipeline = sd_pipeline270 271        self.sd_pipeline.set_progress_bar_config(272            position=1,273            disable=not self.accelerator.is_local_main_process,274            leave=False,275            desc="Timestep",276            dynamic_ncols=True,277        )278 279        # For mixed precision training we cast all non-trainable weights (vae, non-lora text_encoder and non-lora unet) to half-precision280        # as these weights are only used for inference, keeping weights in full precision is not required.281        if self.accelerator.mixed_precision == "fp16":282            inference_dtype = torch.float16283        elif self.accelerator.mixed_precision == "bf16":284            inference_dtype = torch.bfloat16285        else:286            inference_dtype = torch.float32287 288        self.sd_pipeline.vae.to(self.accelerator.device, dtype=inference_dtype)289        self.sd_pipeline.text_encoder.to(self.accelerator.device, dtype=inference_dtype)290        self.sd_pipeline.unet.to(self.accelerator.device, dtype=inference_dtype)291 292        trainable_layers = self.sd_pipeline.get_trainable_layers()293 294        self.accelerator.register_save_state_pre_hook(self._save_model_hook)295        self.accelerator.register_load_state_pre_hook(self._load_model_hook)296 297        # Enable TF32 for faster training on Ampere GPUs,298        # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices299        if self.config.allow_tf32:300            torch.backends.cuda.matmul.allow_tf32 = True301 302        self.optimizer = self._setup_optimizer(303            trainable_layers.parameters() if not isinstance(trainable_layers, list) else trainable_layers304        )305 306        self.neg_prompt_embed = self.sd_pipeline.text_encoder(307            self.sd_pipeline.tokenizer(308                [""] if self.config.negative_prompts is None else self.config.negative_prompts,309                return_tensors="pt",310                padding="max_length",311                truncation=True,312                max_length=self.sd_pipeline.tokenizer.model_max_length,313            ).input_ids.to(self.accelerator.device)314        )[0]315 316        # NOTE: for some reason, autocast is necessary for non-lora training but for lora training it isn't necessary and it uses317        # more memory318        self.autocast = self.sd_pipeline.autocast or self.accelerator.autocast319 320        if hasattr(self.sd_pipeline, "use_lora") and self.sd_pipeline.use_lora:321            unet, self.optimizer = self.accelerator.prepare(trainable_layers, self.optimizer)322            self.trainable_layers = list(filter(lambda p: p.requires_grad, unet.parameters()))323        else:324            self.trainable_layers, self.optimizer = self.accelerator.prepare(trainable_layers, self.optimizer)325 326        if config.resume_from:327            logger.info(f"Resuming from {config.resume_from}")328            self.accelerator.load_state(config.resume_from)329            self.first_epoch = int(config.resume_from.split("_")[-1]) + 1330        else:331            self.first_epoch = 0332 333    def compute_rewards(self, prompt_image_pairs):334        reward, reward_metadata = self.reward_fn(335            prompt_image_pairs["images"], prompt_image_pairs["prompts"], prompt_image_pairs["prompt_metadata"]336        )337        return reward338 339    def step(self, epoch: int, global_step: int):340        """341        Perform a single step of training.342 343        Args:344            epoch (int): The current epoch.345            global_step (int): The current global step.346 347        Side Effects:348            - Model weights are updated349            - Logs the statistics to the accelerator trackers.350            - If `self.image_samples_callback` is not None, it will be called with the prompt_image_pairs, global_step, and the accelerator tracker.351 352        Returns:353            global_step (int): The updated global step.354        """355        info = defaultdict(list)356 357        self.sd_pipeline.unet.train()358 359        for _ in range(self.config.train_gradient_accumulation_steps):360            with self.accelerator.accumulate(self.sd_pipeline.unet), self.autocast(), torch.enable_grad():361                prompt_image_pairs = self._generate_samples(362                    batch_size=self.config.train_batch_size,363                )364 365                rewards = self.compute_rewards(prompt_image_pairs)366 367                prompt_image_pairs["rewards"] = rewards368 369                rewards_vis = self.accelerator.gather(rewards).detach().cpu().numpy()370 371                loss = self.calculate_loss(rewards)372 373                self.accelerator.backward(loss)374 375                if self.accelerator.sync_gradients:376                    self.accelerator.clip_grad_norm_(377                        self.trainable_layers.parameters()378                        if not isinstance(self.trainable_layers, list)379                        else self.trainable_layers,380                        self.config.train_max_grad_norm,381                    )382 383                self.optimizer.step()384                self.optimizer.zero_grad()385 386            info["reward_mean"].append(rewards_vis.mean())387            info["reward_std"].append(rewards_vis.std())388            info["loss"].append(loss.item())389 390        # Checks if the accelerator has performed an optimization step behind the scenes391        if self.accelerator.sync_gradients:392            # log training-related stuff393            info = {k: torch.mean(torch.tensor(v)) for k, v in info.items()}394            info = self.accelerator.reduce(info, reduction="mean")395            info.update({"epoch": epoch})396            self.accelerator.log(info, step=global_step)397            global_step += 1398            info = defaultdict(list)399        else:400            raise ValueError(401                "Optimization step should have been performed by this point. Please check calculated gradient accumulation settings."402            )403        # Logs generated images404        if self.image_samples_callback is not None and global_step % self.config.log_image_freq == 0:405            self.image_samples_callback(prompt_image_pairs, global_step, self.accelerator.trackers[0])406 407        if epoch != 0 and epoch % self.config.save_freq == 0 and self.accelerator.is_main_process:408            self.accelerator.save_state()409 410        return global_step411 412    def calculate_loss(self, rewards):413        """414        Calculate the loss for a batch of an unpacked sample415 416        Args:417            rewards (torch.Tensor):418                Differentiable reward scalars for each generated image, shape: [batch_size]419 420        Returns:421            loss (torch.Tensor)422            (all of these are of shape (1,))423        """424        #  Loss is specific to Aesthetic Reward function used in AlignProp (https://huggingface.co/papers/2310.03739)425        loss = 10.0 - (rewards).mean()426        return loss427 428    def loss(429        self,430        advantages: torch.Tensor,431        clip_range: float,432        ratio: torch.Tensor,433    ):434        unclipped_loss = -advantages * ratio435        clipped_loss = -advantages * torch.clamp(436            ratio,437            1.0 - clip_range,438            1.0 + clip_range,439        )440        return torch.mean(torch.maximum(unclipped_loss, clipped_loss))441 442    def _setup_optimizer(self, trainable_layers_parameters):443        if self.config.train_use_8bit_adam:444            import bitsandbytes445 446            optimizer_cls = bitsandbytes.optim.AdamW8bit447        else:448            optimizer_cls = torch.optim.AdamW449 450        return optimizer_cls(451            trainable_layers_parameters,452            lr=self.config.train_learning_rate,453            betas=(self.config.train_adam_beta1, self.config.train_adam_beta2),454            weight_decay=self.config.train_adam_weight_decay,455            eps=self.config.train_adam_epsilon,456        )457 458    def _save_model_hook(self, models, weights, output_dir):459        self.sd_pipeline.save_checkpoint(models, weights, output_dir)460        weights.pop()  # ensures that accelerate doesn't try to handle saving of the model461 462    def _load_model_hook(self, models, input_dir):463        self.sd_pipeline.load_checkpoint(models, input_dir)464        models.pop()  # ensures that accelerate doesn't try to handle loading of the model465 466    def _generate_samples(self, batch_size, with_grad=True, prompts=None):467        """468        Generate samples from the model469 470        Args:471            batch_size (int): Batch size to use for sampling472            with_grad (bool): Whether the generated RGBs should have gradients attached to it.473 474        Returns:475            prompt_image_pairs (dict[Any])476        """477        prompt_image_pairs = {}478 479        sample_neg_prompt_embeds = self.neg_prompt_embed.repeat(batch_size, 1, 1)480 481        if prompts is None:482            prompts, prompt_metadata = zip(*[self.prompt_fn() for _ in range(batch_size)])483        else:484            prompt_metadata = [{} for _ in range(batch_size)]485 486        prompt_ids = self.sd_pipeline.tokenizer(487            prompts,488            return_tensors="pt",489            padding="max_length",490            truncation=True,491            max_length=self.sd_pipeline.tokenizer.model_max_length,492        ).input_ids.to(self.accelerator.device)493 494        prompt_embeds = self.sd_pipeline.text_encoder(prompt_ids)[0]495 496        if with_grad:497            sd_output = self.sd_pipeline.rgb_with_grad(498                prompt_embeds=prompt_embeds,499                negative_prompt_embeds=sample_neg_prompt_embeds,500                num_inference_steps=self.config.sample_num_steps,501                guidance_scale=self.config.sample_guidance_scale,502                eta=self.config.sample_eta,503                truncated_backprop_rand=self.config.truncated_backprop_rand,504                truncated_backprop_timestep=self.config.truncated_backprop_timestep,505                truncated_rand_backprop_minmax=self.config.truncated_rand_backprop_minmax,506                output_type="pt",507            )508        else:509            sd_output = self.sd_pipeline(510                prompt_embeds=prompt_embeds,511                negative_prompt_embeds=sample_neg_prompt_embeds,512                num_inference_steps=self.config.sample_num_steps,513                guidance_scale=self.config.sample_guidance_scale,514                eta=self.config.sample_eta,515                output_type="pt",516            )517 518        images = sd_output.images519 520        prompt_image_pairs["images"] = images521        prompt_image_pairs["prompts"] = prompts522        prompt_image_pairs["prompt_metadata"] = prompt_metadata523 524        return prompt_image_pairs525 526    def train(self, epochs: Optional[int] = None):527        """528        Train the model for a given number of epochs529        """530        global_step = 0531        if epochs is None:532            epochs = self.config.num_epochs533        for epoch in range(self.first_epoch, epochs):534            global_step = self.step(epoch, global_step)535 536    def _save_pretrained(self, save_directory):537        self.sd_pipeline.save_pretrained(save_directory)538        self.create_model_card()539 540    def create_model_card(541        self,542        model_name: Optional[str] = None,543        dataset_name: Optional[str] = None,544        tags: Union[str, list[str], None] = None,545    ):546        """547        Creates a draft of a model card using the information available to the `Trainer`.548 549        Args:550            model_name (`str` or `None`, *optional*, defaults to `None`):551                Name of the model.552            dataset_name (`str` or `None`, *optional*, defaults to `None`):553                Name of the dataset used for training.554            tags (`str`, `list[str]` or `None`, *optional*, defaults to `None`):555                Tags to be associated with the model card.556        """557        if not self.is_world_process_zero():558            return559 560        if hasattr(self.model.config, "_name_or_path") and not os.path.isdir(self.model.config._name_or_path):561            base_model = self.model.config._name_or_path562        else:563            base_model = None564 565        tags = tags or []566        if isinstance(tags, str):567            tags = [tags]568 569        if hasattr(self.model.config, "unsloth_version"):570            tags.append("unsloth")571 572        citation = textwrap.dedent("""\573        @article{prabhudesai2024aligning,574            title        = {{Aligning Text-to-Image Diffusion Models with Reward Backpropagation}},575            author       = {Mihir Prabhudesai and Anirudh Goyal and Deepak Pathak and Katerina Fragkiadaki},576            year         = 2024,577            eprint       = {arXiv:2310.03739}578        }""")579 580        model_card = generate_model_card(581            base_model=base_model,582            model_name=model_name,583            hub_model_id=self.hub_model_id,584            dataset_name=dataset_name,585            tags=tags,586            wandb_url=wandb.run.get_url() if is_wandb_available() and wandb.run is not None else None,587            comet_url=get_comet_experiment_url(),588            trainer_name="AlignProp",589            trainer_citation=citation,590            paper_title="Aligning Text-to-Image Diffusion Models with Reward Backpropagation",591            paper_id="2310.03739",592        )593 594        model_card.save(os.path.join(self.args.output_dir, "README.md"))595class UnslothAlignPropTrainer(_UnslothAlignPropTrainer):596    """597    598    The AlignPropTrainer uses Deep Diffusion Policy Optimization to optimise diffusion models.599    Note, this trainer is heavily inspired by the work here: https://github.com/mihirp1998/AlignProp/600    As of now only Stable Diffusion based pipelines are supported601 602    Attributes:603        config (`AlignPropConfig`):604            Configuration object for AlignPropTrainer. Check the documentation of `PPOConfig` for more details.605        reward_function (`Callable[[torch.Tensor, tuple[str], tuple[Any]], torch.Tensor]`):606            Reward function to be used607        prompt_function (`Callable[[], tuple[str, Any]]`):608            Function to generate prompts to guide model609        sd_pipeline (`DDPOStableDiffusionPipeline`):610            Stable Diffusion pipeline to be used for training.611        image_samples_hook (`Optional[Callable[[Any, Any, Any], Any]]`):612            Hook to be called to log images613    614    """615    def __init__(616        self,617        config,618        reward_function,619        prompt_function,620        sd_pipeline,621        image_samples_hook = None,622        **kwargs623    ):624        if args is None: args = UnslothAlignPropConfig()625        other_metrics = []626        627        from unsloth_zoo.logging_utils import PatchRLStatistics628        PatchRLStatistics('alignprop_trainer', other_metrics)629        630        super().__init__(631            config = config,632            reward_function = reward_function,633            prompt_function = prompt_function,634            sd_pipeline = sd_pipeline,635            image_samples_hook = image_samples_hook,**kwargs)636        637pass638