CoolFace
Apppublic

lorocksUMD/LlavaMistral1

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
train_LoHa.py994 linesDownload Raw Back to finetuning_scripts
1# Adopted from https://github.com/lm-sys/FastChat. Below is the original copyright:2# Adopted from tatsu-lab@stanford_alpaca. Below is the original copyright:3#    Copyright 2023 Rohan Taori, Ishaan Gulrajani, Tianyi Zhang, Yann Dubois, Xuechen Li4#5#    Licensed under the Apache License, Version 2.0 (the "License");6#    you may not use this file except in compliance with the License.7#    You may obtain a copy of the License at8#9#        http://www.apache.org/licenses/LICENSE-2.010#11#    Unless required by applicable law or agreed to in writing, software12#    distributed under the License is distributed on an "AS IS" BASIS,13#    WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.14#    See the License for the specific language governing permissions and15#    limitations under the License.16 17# Minor edits by Lowell Lobo18 19import os20import copy21from dataclasses import dataclass, field22import json23import logging24import pathlib25from typing import Dict, Optional, Sequence, List26 27import torch28 29import transformers30import tokenizers31 32from llava.constants import IGNORE_INDEX, IMAGE_TOKEN_INDEX, DEFAULT_IMAGE_TOKEN, DEFAULT_IM_START_TOKEN, DEFAULT_IM_END_TOKEN33from torch.utils.data import Dataset34from llava.train.llava_trainer import LLaVATrainer35 36from llava import conversation as conversation_lib37from llava.model import *38from llava.mm_utils import tokenizer_image_token39 40from PIL import Image41 42 43local_rank = None44 45 46def rank0_print(*args):47    if local_rank == 0:48        print(*args)49 50 51from packaging import version52IS_TOKENIZER_GREATER_THAN_0_14 = version.parse(tokenizers.__version__) >= version.parse('0.14')53 54 55@dataclass56class ModelArguments:57    model_name_or_path: Optional[str] = field(default="facebook/opt-125m")58    version: Optional[str] = field(default="v0")59    freeze_backbone: bool = field(default=False)60    tune_mm_mlp_adapter: bool = field(default=False)61    vision_tower: Optional[str] = field(default=None)62    mm_vision_select_layer: Optional[int] = field(default=-1)   # default to the last layer63    pretrain_mm_mlp_adapter: Optional[str] = field(default=None)64    mm_projector_type: Optional[str] = field(default='linear')65    mm_use_im_start_end: bool = field(default=False)66    mm_use_im_patch_token: bool = field(default=True)67    mm_patch_merge_type: Optional[str] = field(default='flat')68    mm_vision_select_feature: Optional[str] = field(default="patch")69 70 71@dataclass72class DataArguments:73    data_path: str = field(default=None,74                           metadata={"help": "Path to the training data."})75    lazy_preprocess: bool = False76    is_multimodal: bool = False77    image_folder: Optional[str] = field(default=None)78    image_aspect_ratio: str = 'square'79 80 81@dataclass82class TrainingArguments(transformers.TrainingArguments):83    cache_dir: Optional[str] = field(default=None)84    optim: str = field(default="adamw_torch")85    remove_unused_columns: bool = field(default=False)86    freeze_mm_mlp_adapter: bool = field(default=False)87    mpt_attn_impl: Optional[str] = field(default="triton")88    model_max_length: int = field(89        default=512,90        metadata={91            "help":92            "Maximum sequence length. Sequences will be right padded (and possibly truncated)."93        },94    )95    double_quant: bool = field(96        default=True,97        metadata={"help": "Compress the quantization statistics through double quantization."}98    )99    quant_type: str = field(100        default="nf4",101        metadata={"help": "Quantization data type to use. Should be one of `fp4` or `nf4`."}102    )103    bits: int = field(104        default=16,105        metadata={"help": "How many bits to use."}106    )107    lora_enable: bool = False108    lora_r: int = 64109    lora_alpha: int = 16110    lora_dropout: float = 0.05111    lora_weight_path: str = ""112    lora_bias: str = "none"113    mm_projector_lr: Optional[float] = None114    group_by_modality_length: bool = field(default=False)115 116 117def maybe_zero_3(param, ignore_status=False, name=None):118    from deepspeed import zero119    from deepspeed.runtime.zero.partition_parameters import ZeroParamStatus120    if hasattr(param, "ds_id"):121        if param.ds_status == ZeroParamStatus.NOT_AVAILABLE:122            if not ignore_status:123                logging.warning(f"{name}: param.ds_status != ZeroParamStatus.NOT_AVAILABLE: {param.ds_status}")124        with zero.GatheredParameters([param]):125            param = param.data.detach().cpu().clone()126    else:127        param = param.detach().cpu().clone()128    return param129 130 131# Borrowed from peft.utils.get_peft_model_state_dict132def get_peft_state_maybe_zero_3(named_params, bias):133    if bias == "none":134        to_return = {k: t for k, t in named_params if "lora_" in k}135    elif bias == "all":136        to_return = {k: t for k, t in named_params if "lora_" in k or "bias" in k}137    elif bias == "lora_only":138        to_return = {}139        maybe_lora_bias = {}140        lora_bias_names = set()141        for k, t in named_params:142            if "lora_" in k:143                to_return[k] = t144                bias_name = k.split("lora_")[0] + "bias"145                lora_bias_names.add(bias_name)146            elif "bias" in k:147                maybe_lora_bias[k] = t148        for k, t in maybe_lora_bias:149            if bias_name in lora_bias_names:150                to_return[bias_name] = t151    else:152        raise NotImplementedError153    to_return = {k: maybe_zero_3(v, ignore_status=True) for k, v in to_return.items()}154    return to_return155 156 157def get_peft_state_non_lora_maybe_zero_3(named_params, require_grad_only=True):158    to_return = {k: t for k, t in named_params if "lora_" not in k}159    if require_grad_only:160        to_return = {k: t for k, t in to_return.items() if t.requires_grad}161    to_return = {k: maybe_zero_3(v, ignore_status=True).cpu() for k, v in to_return.items()}162    return to_return163 164 165def get_mm_adapter_state_maybe_zero_3(named_params, keys_to_match):166    to_return = {k: t for k, t in named_params if any(key_match in k for key_match in keys_to_match)}167    to_return = {k: maybe_zero_3(v, ignore_status=True).cpu() for k, v in to_return.items()}168    return to_return169 170 171def find_all_linear_names(model):172    cls = torch.nn.Linear173    lora_module_names = set()174    multimodal_keywords = ['mm_projector', 'vision_tower', 'vision_resampler']175    for name, module in model.named_modules():176        if any(mm_keyword in name for mm_keyword in multimodal_keywords):177            continue178        if isinstance(module, cls):179            names = name.split('.')180            lora_module_names.add(names[0] if len(names) == 1 else names[-1])181 182    if 'lm_head' in lora_module_names: # needed for 16-bit183        lora_module_names.remove('lm_head')184    return list(lora_module_names)185 186 187def safe_save_model_for_hf_trainer(trainer: transformers.Trainer,188                                   output_dir: str):189    """Collects the state dict and dump to disk."""190 191    if getattr(trainer.args, "tune_mm_mlp_adapter", False):192        # Only save Adapter193        keys_to_match = ['mm_projector']194        if getattr(trainer.args, "use_im_start_end", False):195            keys_to_match.extend(['embed_tokens', 'embed_in'])196 197        weight_to_save = get_mm_adapter_state_maybe_zero_3(trainer.model.named_parameters(), keys_to_match)198        trainer.model.config.save_pretrained(output_dir)199 200        current_folder = output_dir.split('/')[-1]201        parent_folder = os.path.dirname(output_dir)202        if trainer.args.local_rank == 0 or trainer.args.local_rank == -1:203            if current_folder.startswith('checkpoint-'):204                mm_projector_folder = os.path.join(parent_folder, "mm_projector")205                os.makedirs(mm_projector_folder, exist_ok=True)206                torch.save(weight_to_save, os.path.join(mm_projector_folder, f'{current_folder}.bin'))207            else:208                torch.save(weight_to_save, os.path.join(output_dir, f'mm_projector.bin'))209        return210 211    if trainer.deepspeed:212        torch.cuda.synchronize()213        trainer.save_model(output_dir)214        return215 216    state_dict = trainer.model.state_dict()217    if trainer.args.should_save:218        cpu_state_dict = {219            key: value.cpu()220            for key, value in state_dict.items()221        }222        del state_dict223        trainer._save(output_dir, state_dict=cpu_state_dict)  # noqa224 225 226def smart_tokenizer_and_embedding_resize(227    special_tokens_dict: Dict,228    tokenizer: transformers.PreTrainedTokenizer,229    model: transformers.PreTrainedModel,230):231    """Resize tokenizer and embedding.232 233    Note: This is the unoptimized version that may make your embedding size not be divisible by 64.234    """235    num_new_tokens = tokenizer.add_special_tokens(special_tokens_dict)236    model.resize_token_embeddings(len(tokenizer))237 238    if num_new_tokens > 0:239        input_embeddings = model.get_input_embeddings().weight.data240        output_embeddings = model.get_output_embeddings().weight.data241 242        input_embeddings_avg = input_embeddings[:-num_new_tokens].mean(243            dim=0, keepdim=True)244        output_embeddings_avg = output_embeddings[:-num_new_tokens].mean(245            dim=0, keepdim=True)246 247        input_embeddings[-num_new_tokens:] = input_embeddings_avg248        output_embeddings[-num_new_tokens:] = output_embeddings_avg249 250 251def _tokenize_fn(strings: Sequence[str],252                 tokenizer: transformers.PreTrainedTokenizer) -> Dict:253    """Tokenize a list of strings."""254    tokenized_list = [255        tokenizer(256            text,257            return_tensors="pt",258            padding="longest",259            max_length=tokenizer.model_max_length,260            truncation=True,261        ) for text in strings262    ]263    input_ids = labels = [264        tokenized.input_ids[0] for tokenized in tokenized_list265    ]266    input_ids_lens = labels_lens = [267        tokenized.input_ids.ne(tokenizer.pad_token_id).sum().item()268        for tokenized in tokenized_list269    ]270    return dict(271        input_ids=input_ids,272        labels=labels,273        input_ids_lens=input_ids_lens,274        labels_lens=labels_lens,275    )276 277 278def _mask_targets(target, tokenized_lens, speakers):279    # cur_idx = 0280    cur_idx = tokenized_lens[0]281    tokenized_lens = tokenized_lens[1:]282    target[:cur_idx] = IGNORE_INDEX283    for tokenized_len, speaker in zip(tokenized_lens, speakers):284        if speaker == "human":285            target[cur_idx+2:cur_idx + tokenized_len] = IGNORE_INDEX286        cur_idx += tokenized_len287 288 289def _add_speaker_and_signal(header, source, get_conversation=True):290    """Add speaker and start/end signal on each round."""291    BEGIN_SIGNAL = "### "292    END_SIGNAL = "\n"293    conversation = header294    for sentence in source:295        from_str = sentence["from"]296        if from_str.lower() == "human":297            from_str = conversation_lib.default_conversation.roles[0]298        elif from_str.lower() == "gpt":299            from_str = conversation_lib.default_conversation.roles[1]300        else:301            from_str = 'unknown'302        sentence["value"] = (BEGIN_SIGNAL + from_str + ": " +303                             sentence["value"] + END_SIGNAL)304        if get_conversation:305            conversation += sentence["value"]306    conversation += BEGIN_SIGNAL307    return conversation308 309 310def preprocess_multimodal(311    sources: Sequence[str],312    data_args: DataArguments313) -> Dict:314    is_multimodal = data_args.is_multimodal315    if not is_multimodal:316        return sources317 318    for source in sources:319        for sentence in source:320            if DEFAULT_IMAGE_TOKEN in sentence['value']:321                sentence['value'] = sentence['value'].replace(DEFAULT_IMAGE_TOKEN, '').strip()322                sentence['value'] = DEFAULT_IMAGE_TOKEN + '\n' + sentence['value']323                sentence['value'] = sentence['value'].strip()324                if "mmtag" in conversation_lib.default_conversation.version:325                    sentence['value'] = sentence['value'].replace(DEFAULT_IMAGE_TOKEN, '<Image>' + DEFAULT_IMAGE_TOKEN + '</Image>')326            replace_token = DEFAULT_IMAGE_TOKEN327            if data_args.mm_use_im_start_end:328                replace_token = DEFAULT_IM_START_TOKEN + replace_token + DEFAULT_IM_END_TOKEN329            sentence["value"] = sentence["value"].replace(DEFAULT_IMAGE_TOKEN, replace_token)330 331    return sources332 333 334def preprocess_llama_2(335    sources,336    tokenizer: transformers.PreTrainedTokenizer,337    has_image: bool = False338) -> Dict:339    conv = conversation_lib.default_conversation.copy()340    roles = {"human": conv.roles[0], "gpt": conv.roles[1]}341 342    # Apply prompt templates343    conversations = []344    for i, source in enumerate(sources):345        if roles[source[0]["from"]] != conv.roles[0]:346            # Skip the first one if it is not from human347            source = source[1:]348 349        conv.messages = []350        for j, sentence in enumerate(source):351            role = roles[sentence["from"]]352            assert role == conv.roles[j % 2], f"{i}"353            conv.append_message(role, sentence["value"])354        conversations.append(conv.get_prompt())355 356    # Tokenize conversations357 358    if has_image:359        input_ids = torch.stack([tokenizer_image_token(prompt, tokenizer, return_tensors='pt') for prompt in conversations], dim=0)360    else:361        input_ids = tokenizer(362            conversations,363            return_tensors="pt",364            padding="longest",365            max_length=tokenizer.model_max_length,366            truncation=True,367        ).input_ids368 369    targets = input_ids.clone()370 371    assert conv.sep_style == conversation_lib.SeparatorStyle.LLAMA_2372 373    # Mask targets374    sep = "[/INST] "375    for conversation, target in zip(conversations, targets):376        total_len = int(target.ne(tokenizer.pad_token_id).sum())377 378        rounds = conversation.split(conv.sep2)379        cur_len = 1380        target[:cur_len] = IGNORE_INDEX381        for i, rou in enumerate(rounds):382            if rou == "":383                break384 385            parts = rou.split(sep)386            if len(parts) != 2:387                break388            parts[0] += sep389 390            if has_image:391                round_len = len(tokenizer_image_token(rou, tokenizer))392                instruction_len = len(tokenizer_image_token(parts[0], tokenizer)) - 2393            else:394                round_len = len(tokenizer(rou).input_ids)395                instruction_len = len(tokenizer(parts[0]).input_ids) - 2396 397            target[cur_len : cur_len + instruction_len] = IGNORE_INDEX398 399            cur_len += round_len400        target[cur_len:] = IGNORE_INDEX401 402        if cur_len < tokenizer.model_max_length:403            if cur_len != total_len:404                target[:] = IGNORE_INDEX405                print(406                    f"WARNING: tokenization mismatch: {cur_len} vs. {total_len}."407                    f" (ignored)"408                )409 410    return dict(411        input_ids=input_ids,412        labels=targets,413    )414 415 416def preprocess_v1(417    sources,418    tokenizer: transformers.PreTrainedTokenizer,419    has_image: bool = False420) -> Dict:421    conv = conversation_lib.default_conversation.copy()422    roles = {"human": conv.roles[0], "gpt": conv.roles[1]}423 424    # Apply prompt templates425    conversations = []426    for i, source in enumerate(sources):427        if roles[source[0]["from"]] != conv.roles[0]:428            # Skip the first one if it is not from human429            source = source[1:]430 431        conv.messages = []432        for j, sentence in enumerate(source):433            role = roles[sentence["from"]]434            assert role == conv.roles[j % 2], f"{i}"435            conv.append_message(role, sentence["value"])436        conversations.append(conv.get_prompt())437 438    # Tokenize conversations439 440    if has_image:441        input_ids = torch.stack([tokenizer_image_token(prompt, tokenizer, return_tensors='pt') for prompt in conversations], dim=0)442    else:443        input_ids = tokenizer(444            conversations,445            return_tensors="pt",446            padding="longest",447            max_length=tokenizer.model_max_length,448            truncation=True,449        ).input_ids450 451    targets = input_ids.clone()452 453    assert conv.sep_style == conversation_lib.SeparatorStyle.TWO454 455    # Mask targets456    sep = conv.sep + conv.roles[1] + ": "457    for conversation, target in zip(conversations, targets):458        total_len = int(target.ne(tokenizer.pad_token_id).sum())459 460        rounds = conversation.split(conv.sep2)461        cur_len = 1462        target[:cur_len] = IGNORE_INDEX463        for i, rou in enumerate(rounds):464            if rou == "":465                break466 467            parts = rou.split(sep)468            if len(parts) != 2:469                break470            parts[0] += sep471 472            if has_image:473                round_len = len(tokenizer_image_token(rou, tokenizer))474                instruction_len = len(tokenizer_image_token(parts[0], tokenizer)) - 2475            else:476                round_len = len(tokenizer(rou).input_ids)477                instruction_len = len(tokenizer(parts[0]).input_ids) - 2478 479            if i != 0 and not tokenizer.legacy and IS_TOKENIZER_GREATER_THAN_0_14:480                round_len -= 1481                instruction_len -= 1482 483            target[cur_len : cur_len + instruction_len] = IGNORE_INDEX484 485            cur_len += round_len486        target[cur_len:] = IGNORE_INDEX487 488        if cur_len < tokenizer.model_max_length:489            if cur_len != total_len:490                target[:] = IGNORE_INDEX491                print(492                    f"WARNING: tokenization mismatch: {cur_len} vs. {total_len}."493                    f" (ignored)"494                )495 496    return dict(497        input_ids=input_ids,498        labels=targets,499    )500 501 502def preprocess_mpt(503    sources,504    tokenizer: transformers.PreTrainedTokenizer,505    has_image: bool = False506) -> Dict:507    conv = conversation_lib.default_conversation.copy()508    roles = {"human": conv.roles[0], "gpt": conv.roles[1]}509 510    # Apply prompt templates511    conversations = []512    for i, source in enumerate(sources):513        if roles[source[0]["from"]] != conv.roles[0]:514            # Skip the first one if it is not from human515            source = source[1:]516 517        conv.messages = []518        for j, sentence in enumerate(source):519            role = roles[sentence["from"]]520            assert role == conv.roles[j % 2], f"{i}"521            conv.append_message(role, sentence["value"])522        conversations.append(conv.get_prompt())523 524    # Tokenize conversations525 526    if has_image:527        input_ids = torch.stack([tokenizer_image_token(prompt, tokenizer, return_tensors='pt') for prompt in conversations], dim=0)528    else:529        input_ids = tokenizer(530            conversations,531            return_tensors="pt",532            padding="longest",533            max_length=tokenizer.model_max_length,534            truncation=True,535        ).input_ids536 537    targets = input_ids.clone()538    assert conv.sep_style == conversation_lib.SeparatorStyle.MPT539 540    # Mask targets541    sep = conv.sep + conv.roles[1]542    for conversation, target in zip(conversations, targets):543        total_len = int(target.ne(tokenizer.pad_token_id).sum())544 545        rounds = conversation.split(conv.sep)546        re_rounds = [conv.sep.join(rounds[:3])] # system + user + gpt547        for conv_idx in range(3, len(rounds), 2):548            re_rounds.append(conv.sep.join(rounds[conv_idx:conv_idx+2]))    # user + gpt549        cur_len = 0550        target[:cur_len] = IGNORE_INDEX551        for i, rou in enumerate(re_rounds):552            if rou == "":553                break554 555            parts = rou.split(sep)556            if len(parts) != 2:557                break558            parts[0] += sep559 560            if has_image:561                round_len = len(tokenizer_image_token(rou, tokenizer))562                instruction_len = len(tokenizer_image_token(parts[0], tokenizer)) - 1563            else:564                round_len = len(tokenizer(rou).input_ids)565                instruction_len = len(tokenizer(parts[0]).input_ids) - 1566 567            if i != 0 and getattr(tokenizer, 'legacy', False) and IS_TOKENIZER_GREATER_THAN_0_14:568                round_len += 1569                instruction_len += 1570 571            target[cur_len : cur_len + instruction_len] = IGNORE_INDEX572 573            cur_len += round_len574        target[cur_len:] = IGNORE_INDEX575 576        if cur_len < tokenizer.model_max_length:577            if cur_len != total_len:578                target[:] = IGNORE_INDEX579                print(580                    f"WARNING: tokenization mismatch: {cur_len} vs. {total_len}."581                    f" (ignored)"582                )583 584    return dict(585        input_ids=input_ids,586        labels=targets,587    )588 589 590def preprocess_plain(591    sources: Sequence[str],592    tokenizer: transformers.PreTrainedTokenizer,593) -> Dict:594    # add end signal and concatenate together595    conversations = []596    for source in sources:597        assert len(source) == 2598        assert DEFAULT_IMAGE_TOKEN in source[0]['value']599        source[0]['value'] = DEFAULT_IMAGE_TOKEN600        conversation = source[0]['value'] + source[1]['value'] + conversation_lib.default_conversation.sep601        conversations.append(conversation)602    # tokenize conversations603    input_ids = [tokenizer_image_token(prompt, tokenizer, return_tensors='pt') for prompt in conversations]604    targets = copy.deepcopy(input_ids)605    for target, source in zip(targets, sources):606        tokenized_len = len(tokenizer_image_token(source[0]['value'], tokenizer))607        target[:tokenized_len] = IGNORE_INDEX608 609    return dict(input_ids=input_ids, labels=targets)610 611 612def preprocess(613    sources: Sequence[str],614    tokenizer: transformers.PreTrainedTokenizer,615    has_image: bool = False616) -> Dict:617    """618    Given a list of sources, each is a conversation list. This transform:619    1. Add signal '### ' at the beginning each sentence, with end signal '\n';620    2. Concatenate conversations together;621    3. Tokenize the concatenated conversation;622    4. Make a deepcopy as the target. Mask human words with IGNORE_INDEX.623    """624    if conversation_lib.default_conversation.sep_style == conversation_lib.SeparatorStyle.PLAIN:625        return preprocess_plain(sources, tokenizer)626    if conversation_lib.default_conversation.sep_style == conversation_lib.SeparatorStyle.LLAMA_2:627        return preprocess_llama_2(sources, tokenizer, has_image=has_image)628    if conversation_lib.default_conversation.version.startswith("v1"):629        return preprocess_v1(sources, tokenizer, has_image=has_image)630    if conversation_lib.default_conversation.version == "mpt":631        return preprocess_mpt(sources, tokenizer, has_image=has_image)632    # add end signal and concatenate together633    conversations = []634    for source in sources:635        header = f"{conversation_lib.default_conversation.system}\n\n"636        conversation = _add_speaker_and_signal(header, source)637        conversations.append(conversation)638    # tokenize conversations639    def get_tokenize_len(prompts):640        return [len(tokenizer_image_token(prompt, tokenizer)) for prompt in prompts]641 642    if has_image:643        input_ids = [tokenizer_image_token(prompt, tokenizer, return_tensors='pt') for prompt in conversations]644    else:645        conversations_tokenized = _tokenize_fn(conversations, tokenizer)646        input_ids = conversations_tokenized["input_ids"]647 648    targets = copy.deepcopy(input_ids)649    for target, source in zip(targets, sources):650        if has_image:651            tokenized_lens = get_tokenize_len([header] + [s["value"] for s in source])652        else:653            tokenized_lens = _tokenize_fn([header] + [s["value"] for s in source], tokenizer)["input_ids_lens"]654        speakers = [sentence["from"] for sentence in source]655        _mask_targets(target, tokenized_lens, speakers)656 657    return dict(input_ids=input_ids, labels=targets)658 659 660class LazySupervisedDataset(Dataset):661    """Dataset for supervised fine-tuning."""662 663    def __init__(self, data_path: str,664                 tokenizer: transformers.PreTrainedTokenizer,665                 data_args: DataArguments):666        super(LazySupervisedDataset, self).__init__()667        list_data_dict = json.load(open(data_path, "r"))668 669        rank0_print("Formatting inputs...Skip in lazy mode")670        self.tokenizer = tokenizer671        self.list_data_dict = list_data_dict672        self.data_args = data_args673 674    def __len__(self):675        return len(self.list_data_dict)676 677    @property678    def lengths(self):679        length_list = []680        for sample in self.list_data_dict:681            img_tokens = 128 if 'image' in sample else 0682            length_list.append(sum(len(conv['value'].split()) for conv in sample['conversations']) + img_tokens)683        return length_list684 685    @property686    def modality_lengths(self):687        length_list = []688        for sample in self.list_data_dict:689            cur_len = sum(len(conv['value'].split()) for conv in sample['conversations'])690            cur_len = cur_len if 'image' in sample else -cur_len691            length_list.append(cur_len)692        return length_list693 694    def __getitem__(self, i) -> Dict[str, torch.Tensor]:695        sources = self.list_data_dict[i]696        if isinstance(i, int):697            sources = [sources]698        assert len(sources) == 1, "Don't know why it is wrapped to a list"  # FIXME699        if 'image' in sources[0]:700            image_file = self.list_data_dict[i]['image']701            image_folder = self.data_args.image_folder702            processor = self.data_args.image_processor703            image = Image.open(os.path.join(image_folder, image_file)).convert('RGB')704            if self.data_args.image_aspect_ratio == 'pad':705                def expand2square(pil_img, background_color):706                    width, height = pil_img.size707                    if width == height:708                        return pil_img709                    elif width > height:710                        result = Image.new(pil_img.mode, (width, width), background_color)711                        result.paste(pil_img, (0, (width - height) // 2))712                        return result713                    else:714                        result = Image.new(pil_img.mode, (height, height), background_color)715                        result.paste(pil_img, ((height - width) // 2, 0))716                        return result717                image = expand2square(image, tuple(int(x*255) for x in processor.image_mean))718                image = processor.preprocess(image, return_tensors='pt')['pixel_values'][0]719            else:720                image = processor.preprocess(image, return_tensors='pt')['pixel_values'][0]721            sources = preprocess_multimodal(722                copy.deepcopy([e["conversations"] for e in sources]),723                self.data_args)724        else:725            sources = copy.deepcopy([e["conversations"] for e in sources])726        data_dict = preprocess(727            sources,728            self.tokenizer,729            has_image=('image' in self.list_data_dict[i]))730        if isinstance(i, int):731            data_dict = dict(input_ids=data_dict["input_ids"][0],732                             labels=data_dict["labels"][0])733 734        # image exist in the data735        if 'image' in self.list_data_dict[i]:736            data_dict['image'] = image737        elif self.data_args.is_multimodal:738            # image does not exist in the data, but the model is multimodal739            crop_size = self.data_args.image_processor.crop_size740            data_dict['image'] = torch.zeros(3, crop_size['height'], crop_size['width'])741        return data_dict742 743 744@dataclass745class DataCollatorForSupervisedDataset(object):746    """Collate examples for supervised fine-tuning."""747 748    tokenizer: transformers.PreTrainedTokenizer749 750    def __call__(self, instances: Sequence[Dict]) -> Dict[str, torch.Tensor]:751        input_ids, labels = tuple([instance[key] for instance in instances]752                                  for key in ("input_ids", "labels"))753        input_ids = torch.nn.utils.rnn.pad_sequence(754            input_ids,755            batch_first=True,756            padding_value=self.tokenizer.pad_token_id)757        labels = torch.nn.utils.rnn.pad_sequence(labels,758                                                 batch_first=True,759                                                 padding_value=IGNORE_INDEX)760        input_ids = input_ids[:, :self.tokenizer.model_max_length]761        labels = labels[:, :self.tokenizer.model_max_length]762        batch = dict(763            input_ids=input_ids,764            labels=labels,765            attention_mask=input_ids.ne(self.tokenizer.pad_token_id),766        )767 768        if 'image' in instances[0]:769            images = [instance['image'] for instance in instances]770            if all(x is not None and x.shape == images[0].shape for x in images):771                batch['images'] = torch.stack(images)772            else:773                batch['images'] = images774 775        return batch776 777 778def make_supervised_data_module(tokenizer: transformers.PreTrainedTokenizer,779                                data_args) -> Dict:780    """Make dataset and collator for supervised fine-tuning."""781    train_dataset = LazySupervisedDataset(tokenizer=tokenizer,782                                data_path=data_args.data_path,783                                data_args=data_args)784    data_collator = DataCollatorForSupervisedDataset(tokenizer=tokenizer)785    return dict(train_dataset=train_dataset,786                eval_dataset=None,787                data_collator=data_collator)788 789 790def train(attn_implementation=None):791    global local_rank792 793    parser = transformers.HfArgumentParser(794        (ModelArguments, DataArguments, TrainingArguments))795    model_args, data_args, training_args = parser.parse_args_into_dataclasses()796    local_rank = training_args.local_rank797    compute_dtype = (torch.float16 if training_args.fp16 else (torch.bfloat16 if training_args.bf16 else torch.float32))798 799    bnb_model_from_pretrained_args = {}800    if training_args.bits in [4, 8]:801        from transformers import BitsAndBytesConfig802        bnb_model_from_pretrained_args.update(dict(803            device_map={"": training_args.device},804            load_in_4bit=training_args.bits == 4,805            load_in_8bit=training_args.bits == 8,806            quantization_config=BitsAndBytesConfig(807                load_in_4bit=training_args.bits == 4,808                load_in_8bit=training_args.bits == 8,809                llm_int8_skip_modules=["mm_projector"],810                llm_int8_threshold=6.0,811                llm_int8_has_fp16_weight=False,812                bnb_4bit_compute_dtype=compute_dtype,813                bnb_4bit_use_double_quant=training_args.double_quant,814                bnb_4bit_quant_type=training_args.quant_type # {'fp4', 'nf4'}815            )816        ))817 818    if model_args.vision_tower is not None:819        if 'mpt' in model_args.model_name_or_path:820            config = transformers.AutoConfig.from_pretrained(model_args.model_name_or_path, trust_remote_code=True)821            config.attn_config['attn_impl'] = training_args.mpt_attn_impl822            model = LlavaMptForCausalLM.from_pretrained(823                model_args.model_name_or_path,824                config=config,825                cache_dir=training_args.cache_dir,826                **bnb_model_from_pretrained_args827            )828        else:829            model = LlavaLlamaForCausalLM.from_pretrained(830                model_args.model_name_or_path,831                cache_dir=training_args.cache_dir,832                attn_implementation=attn_implementation,833                torch_dtype=(torch.bfloat16 if training_args.bf16 else None),834                **bnb_model_from_pretrained_args835            )836    else:837        model = transformers.LlamaForCausalLM.from_pretrained(838            model_args.model_name_or_path,839            cache_dir=training_args.cache_dir,840            attn_implementation=attn_implementation,841            torch_dtype=(torch.bfloat16 if training_args.bf16 else None),842            **bnb_model_from_pretrained_args843        )844    model.config.use_cache = False845 846    if model_args.freeze_backbone:847        model.model.requires_grad_(False)848 849    if training_args.bits in [4, 8]:850        from peft import prepare_model_for_kbit_training851        model.config.torch_dtype=(torch.float32 if training_args.fp16 else (torch.bfloat16 if training_args.bf16 else torch.float32))852        model = prepare_model_for_kbit_training(model, use_gradient_checkpointing=training_args.gradient_checkpointing)853 854    if training_args.gradient_checkpointing:855        if hasattr(model, "enable_input_require_grads"):856            model.enable_input_require_grads()857        else:858            def make_inputs_require_grad(module, input, output):859                output.requires_grad_(True)860            model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)861 862    if training_args.lora_enable:863        from peft import LoHaConfig, get_peft_model864        lora_config = LoHaConfig(865            r=training_args.lora_r,866            alpha=training_args.lora_alpha,867            target_modules=find_all_linear_names(model),868            rank_dropout=training_args.lora_dropout,869            # module_dropout=training_args.lora_droupout, # New param870            # bias=training_args.lora_bias, # Not there for LoHa hmmmm871            task_type="CAUSAL_LM",872        )873        if training_args.bits == 16:874            if training_args.bf16:875                model.to(torch.bfloat16)876            if training_args.fp16:877                model.to(torch.float16)878        rank0_print("Adding LoRA adapters...")879        model = get_peft_model(model, lora_config)880 881    if 'mpt' in model_args.model_name_or_path:882        tokenizer = transformers.AutoTokenizer.from_pretrained(883            model_args.model_name_or_path,884            cache_dir=training_args.cache_dir,885            model_max_length=training_args.model_max_length,886            padding_side="right"887        )888    else:889        tokenizer = transformers.AutoTokenizer.from_pretrained(890            model_args.model_name_or_path,891            cache_dir=training_args.cache_dir,892            model_max_length=training_args.model_max_length,893            padding_side="right",894            use_fast=False,895        )896 897    if model_args.version == "v0":898        if tokenizer.pad_token is None:899            smart_tokenizer_and_embedding_resize(900                special_tokens_dict=dict(pad_token="[PAD]"),901                tokenizer=tokenizer,902                model=model,903            )904    elif model_args.version == "v0.5":905        tokenizer.pad_token = tokenizer.unk_token906    else:907        tokenizer.pad_token = tokenizer.unk_token908        if model_args.version in conversation_lib.conv_templates:909            conversation_lib.default_conversation = conversation_lib.conv_templates[model_args.version]910        else:911            conversation_lib.default_conversation = conversation_lib.conv_templates["vicuna_v1"]912 913    if model_args.vision_tower is not None:914        model.get_model().initialize_vision_modules(915            model_args=model_args,916            fsdp=training_args.fsdp917        )918        919        vision_tower = model.get_vision_tower()920        vision_tower.to(dtype=torch.bfloat16 if training_args.bf16 else torch.float16, device=training_args.device)921 922        data_args.image_processor = vision_tower.image_processor923        data_args.is_multimodal = True924 925        model.config.image_aspect_ratio = data_args.image_aspect_ratio926        model.config.tokenizer_padding_side = tokenizer.padding_side927        model.config.tokenizer_model_max_length = tokenizer.model_max_length928 929        model.config.tune_mm_mlp_adapter = training_args.tune_mm_mlp_adapter = model_args.tune_mm_mlp_adapter930        if model_args.tune_mm_mlp_adapter:931            model.requires_grad_(False)932            for p in model.get_model().mm_projector.parameters():933                p.requires_grad = True934 935        model.config.freeze_mm_mlp_adapter = training_args.freeze_mm_mlp_adapter936        if training_args.freeze_mm_mlp_adapter:937            for p in model.get_model().mm_projector.parameters():938                p.requires_grad = False939 940        if training_args.bits in [4, 8]:941            model.get_model().mm_projector.to(dtype=compute_dtype, device=training_args.device)942 943        model.config.mm_use_im_start_end = data_args.mm_use_im_start_end = model_args.mm_use_im_start_end944        model.config.mm_projector_lr = training_args.mm_projector_lr945        training_args.use_im_start_end = model_args.mm_use_im_start_end946        model.config.mm_use_im_patch_token = model_args.mm_use_im_patch_token947        model.initialize_vision_tokenizer(model_args, tokenizer=tokenizer)948 949    if training_args.bits in [4, 8]:950        from peft.tuners.loha import LoHaLayer951        for name, module in model.named_modules():952            if isinstance(module, LoHaLayer):953                if training_args.bf16:954                    module = module.to(torch.bfloat16)955            if 'norm' in name:956                module = module.to(torch.float32)957            if 'lm_head' in name or 'embed_tokens' in name:958                if hasattr(module, 'weight'):959                    if training_args.bf16 and module.weight.dtype == torch.float32:960                        module = module.to(torch.bfloat16)961 962    data_module = make_supervised_data_module(tokenizer=tokenizer,963                                              data_args=data_args)964    trainer = LLaVATrainer(model=model,965                    tokenizer=tokenizer,966                    args=training_args,967                    **data_module)968 969    if list(pathlib.Path(training_args.output_dir).glob("checkpoint-*")):970        trainer.train(resume_from_checkpoint=True)971    else:972        trainer.train()973    trainer.save_state()974 975    model.config.use_cache = True976 977    if training_args.lora_enable:978        state_dict = get_peft_state_maybe_zero_3(979            model.named_parameters(), training_args.lora_bias980        )981        non_lora_state_dict = get_peft_state_non_lora_maybe_zero_3(982            model.named_parameters()983        )984        if training_args.local_rank == 0 or training_args.local_rank == -1:985            model.config.save_pretrained(training_args.output_dir)986            model.save_pretrained(training_args.output_dir, state_dict=state_dict)987            torch.save(non_lora_state_dict, os.path.join(training_args.output_dir, 'non_loha_trainables.bin'))988    else:989        safe_save_model_for_hf_trainer(trainer=trainer,990                                       output_dir=training_args.output_dir)991 992 993if __name__ == "__main__":994    train()