Dataset for supervised fine-tuning.
| 21 | llama3_chat_template = "{% set loop_messages = messages %}{% for message in loop_messages %}{% set content = '<|start_header_id|>' + message['role'] + '<|end_header_id|>\n\n'+ message['content'] | trim + '<|eot_id|>' %}{% if loop.index0 == 0 %}{% set content = bos_token + content %}{% endif %}{{ content }}{% endfor %}" |
| 22 | |
| 23 | class SupervisedDataset(Dataset): |
| 24 | """Dataset for supervised fine-tuning.""" |
| 25 | |
| 26 | def __init__( |
| 27 | self, |
| 28 | raw_data, |
| 29 | transform, |
| 30 | tokenizer, |
| 31 | slice_config, |
| 32 | llm_type="minicpm", |
| 33 | patch_size=14, |
| 34 | query_nums=64, |
| 35 | batch_vision=False, |
| 36 | max_length=2048, |
| 37 | max_line_res=1120 |
| 38 | ): |
| 39 | super(SupervisedDataset, self).__init__() |
| 40 | self.raw_data = raw_data |
| 41 | self.tokenizer = tokenizer |
| 42 | self.transform = transform |
| 43 | self.slice_config = slice_config |
| 44 | self.llm_type = llm_type |
| 45 | self.patch_size = patch_size |
| 46 | self.query_nums=query_nums |
| 47 | self.batch_vision = batch_vision |
| 48 | self.max_length = max_length |
| 49 | self.max_line_res = max_line_res |
| 50 | |
| 51 | def __len__(self): |
| 52 | return len(self.raw_data) |
| 53 | |
| 54 | def __resize__(self, origin_img): |
| 55 | resolution = origin_img.size |
| 56 | w,h = resolution |
| 57 | if self.max_line_res is not None: |
| 58 | max_line = self.max_line_res |
| 59 | if h > max_line: |
| 60 | w = int(w * max_line / h) |
| 61 | h = max_line |
| 62 | if w > max_line: |
| 63 | h = int(h * max_line / w) |
| 64 | w = max_line |
| 65 | img = origin_img.resize((w,h),resample=Image.Resampling.LANCZOS) |
| 66 | return img |
| 67 | |
| 68 | def __getitem__(self, i) -> Dict[str, torch.Tensor]: |
| 69 | try: |
| 70 | if isinstance(self.raw_data[i]["image"], str): |
| 71 | # resize the image |
| 72 | images_dict = { "<image>" : self.__resize__(Image.open(self.raw_data[i]["image"]).convert("RGB"))} |
| 73 | elif isinstance(self.raw_data[i]["image"], Dict): |
| 74 | ### for multi-images input, the template for every image is <image_xx>, such as <image_00>, <image_01> |
| 75 | images_dict = {img_name : self.__resize__(Image.open(img_path).convert("RGB")) for img_name, img_path in self.raw_data[i]["image"].items()} |
| 76 | |
| 77 | ret = preprocess( |
| 78 | images_dict, |
| 79 | self.raw_data[i]["conversations"], |
| 80 | self.tokenizer, |
nothing calls this directly
no outgoing calls
no test coverage detected