MCPcopy Create free account
hub / github.com/OpenBMB/AgentCPM-GUI / SupervisedDataset

Class SupervisedDataset

sft/dataset.py:23–101  ·  view source on GitHub ↗

Dataset for supervised fine-tuning.

Source from the content-addressed store, hash-verified

21llama3_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
23class 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,

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected