Setting the file. One moment.
Sam Segmentation Training · Huggingface Vision Trainer · huggingface/skills · Skills Docs
ContentsBack to the top of the page scripts/ sam_segmentation_training.py
Python · 382 lines · 14 KB
17 import math
18 import os
19 import sys
20 from dataclasses import dataclass, field
21 from typing import Any
22
23 import numpy as np
24 import torch
25 import torch.nn.functional as F
26 from datasets import load_dataset
27 from torch.utils.data import Dataset
28
29 import monai
30 import trackio
31
32 import transformers
33 from transformers import (
34 HfArgumentParser,
35 Trainer,
36 TrainingArguments,
37 )
38 from transformers.utils import check_min_version
39
40 logger = logging.getLogger( __name__ )
41
42 check_min_version( "4.57.0.dev0" )
43
44
45 # ---------------------------------------------------------------------------
46 # Dataset wrapper
47 # ---------------------------------------------------------------------------
48
49 class SAMSegmentationDataset ( Dataset ):
50 """Wraps a HF dataset into the format expected by SAM/SAM2 processors.
51
52 Each sample must contain an image, a binary mask, and a prompt (bbox or
53 point). Prompts are read from a JSON-encoded ``prompt`` column or from
54 dedicated ``bbox`` / ``point`` columns.
55 """
56
57 def __init__ (self, dataset, processor, prompt_type: str ,
58 image_col: str , mask_col: str , prompt_col: str | None ,
59 bbox_col: str | None , point_col: str | None ):
60 self .dataset = dataset
61 self .processor = processor
62 self .prompt_type = prompt_type
63 self .image_col = image_col
64 self .mask_col = mask_col
65 self .prompt_col = prompt_col
66 self .bbox_col = bbox_col
67 self .point_col = point_col
68
69 def __len__ (self):
70 return len ( self .dataset)
71
72 def _extract_prompt (self, item):
73 if self .prompt_col and self .prompt_col in item:
74 raw = item[ self .prompt_col]
75 parsed = json.loads(raw) if isinstance (raw, str ) else raw
76 if self .prompt_type == "bbox" :
77 return parsed.get( "bbox" ) or parsed.get( "box" )
78 return parsed.get( "point" ) or parsed.get( "points" )
79
80 if self .prompt_type == "bbox" and self .bbox_col:
81 return item[ self .bbox_col]
82 if self .prompt_type == "point" and self .point_col:
83 return item[ self .point_col]
84 raise ValueError ( "Could not extract prompt from sample" )
85
86 def __getitem__ (self, idx):
87 item = self .dataset[idx]
88 image = item[ self .image_col]
89 prompt = self ._extract_prompt(item)
90
91 if self .prompt_type == "bbox" :
92 inputs = self .processor(image, input_boxes = [[prompt]], return_tensors = "pt" )
93 else :
94 if isinstance (prompt[ 0 ], ( int , float )):
95 prompt = [prompt]
96 inputs = self .processor(image, input_points = [[prompt]], return_tensors = "pt" )
97
98 mask = np.array(item[ self .mask_col])
99 if mask.ndim == 3 :
100 mask = mask[:, :, 0 ]
101 inputs[ "labels" ] = (mask > 0 ).astype(np.float32)
102 inputs[ "original_image_size" ] = torch.tensor(image.size[:: - 1 ])
103 return inputs
104
105
106 def collate_fn (batch):
107 pixel_values = torch.cat([item[ "pixel_values" ] for item in batch], dim = 0 )
108 original_sizes = torch.stack([item[ "original_sizes" ] for item in batch])
109 original_image_size = torch.stack([item[ "original_image_size" ] for item in batch])
110
111 has_boxes = "input_boxes" in batch[ 0 ]
112 has_points = "input_points" in batch[ 0 ]
113
114 labels = torch.cat(
115 [
116 F.interpolate(
117 torch.as_tensor(x[ "labels" ]).unsqueeze( 0 ).unsqueeze( 0 ).float(),
118 size = ( 256 , 256 ),
119 mode = "nearest" ,
120 )
121 for x in batch
122 ],
123 dim = 0 ,
124 ).long()
125
126 result = {
127 "pixel_values" : pixel_values,
128 "original_sizes" : original_sizes,
129 "labels" : labels,
130 "original_image_size" : original_image_size,
131 "multimask_output" : False ,
132 }
133
134 if has_boxes:
135 result[ "input_boxes" ] = torch.cat([item[ "input_boxes" ] for item in batch], dim = 0 )
136 if has_points:
137 result[ "input_points" ] = torch.cat([item[ "input_points" ] for item in batch], dim = 0 )
138 if "input_labels" in batch[ 0 ]:
139 result[ "input_labels" ] = torch.cat([item[ "input_labels" ] for item in batch], dim = 0 )
140
141 return result
142
143
144 # ---------------------------------------------------------------------------
145 # Custom loss (SAM/SAM2 don't compute loss in forward())
146 # ---------------------------------------------------------------------------
147
148 seg_loss = monai.losses.DiceCELoss( sigmoid = True , squared_pred = True , reduction = "mean" )
149
150
151 def compute_loss (outputs, labels, num_items_in_batch = None ):
152 predicted_masks = outputs.pred_masks.squeeze( 1 )
153 return seg_loss(predicted_masks, labels.float())
154
155
156 # ---------------------------------------------------------------------------
157 # CLI arguments
158 # ---------------------------------------------------------------------------
159
160 @dataclass
161 class DataTrainingArguments :
162 dataset_name: str = field(
163 default = "merve/MicroMat-mini" ,
164 metadata = { "help" : "Hub dataset ID." },
165 )
166 dataset_config_name: str | None = field(
167 default = None ,
168 metadata = { "help" : "Dataset config name." },
169 )
170 train_val_split: float | None = field(
171 default = 0.1 ,
172 metadata = { "help" : "Fraction to split off for validation (used when no validation split exists)." },
173 )
174 max_train_samples: int | None = field(
175 default = None ,
176 metadata = { "help" : "Truncate training set (for quick tests)." },
177 )
178 max_eval_samples: int | None = field(
179 default = None ,
180 metadata = { "help" : "Truncate evaluation set." },
181 )
182 image_column_name: str = field(
183 default = "image" ,
184 metadata = { "help" : "Column containing PIL images." },
185 )
186 mask_column_name: str = field(
187 default = "mask" ,
188 metadata = { "help" : "Column containing ground-truth binary masks." },
189 )
190 prompt_column_name: str | None = field(
191 default = "prompt" ,
192 metadata = { "help" : "Column with JSON-encoded prompt (bbox/point). Set to '' to disable." },
193 )
194 bbox_column_name: str | None = field(
195 default = None ,
196 metadata = { "help" : "Column with bbox prompt ([x0,y0,x1,y1]). Used when prompt_column_name is unset." },
197 )
198 point_column_name: str | None = field(
199 default = None ,
200 metadata = { "help" : "Column with point prompt ([x,y] or [[x,y],...]). Used when prompt_column_name is unset." },
201 )
202 prompt_type: str = field(
203 default = "bbox" ,
204 metadata = { "help" : "Prompt type: 'bbox' or 'point'." },
205 )
206
207
208 @dataclass
209 class ModelArguments :
210 model_name_or_path: str = field(
211 default = "facebook/sam2.1-hiera-small" ,
212 metadata = { "help" : "Pretrained SAM/SAM2 model identifier." },
213 )
214 cache_dir: str | None = field( default = None , metadata = { "help" : "Cache directory." })
215 model_revision: str = field( default = "main" , metadata = { "help" : "Model revision." })
216 token: str | None = field( default = None , metadata = { "help" : "Auth token." })
217 trust_remote_code: bool = field( default = False , metadata = { "help" : "Trust remote code." })
218 freeze_vision_encoder: bool = field(
219 default = True ,
220 metadata = { "help" : "Freeze vision encoder weights." },
221 )
222 freeze_prompt_encoder: bool = field(
223 default = True ,
224 metadata = { "help" : "Freeze prompt encoder weights." },
225 )
226
227
228 # ---------------------------------------------------------------------------
229 # Main
230 # ---------------------------------------------------------------------------
231
232 def main ():
233 parser = HfArgumentParser((ModelArguments, DataTrainingArguments, TrainingArguments))
234 parser.set_defaults( per_device_train_batch_size = 4 , num_train_epochs = 30 )
235 if len (sys.argv) == 2 and sys.argv[ 1 ].endswith( ".json" ):
236 model_args, data_args, training_args = parser.parse_json_file(
237 json_file = os.path.abspath(sys.argv[ 1 ])
238 )
239 else :
240 model_args, data_args, training_args = parser.parse_args_into_dataclasses()
241
242 from huggingface_hub import login
243 hf_token = os.environ.get( "HF_TOKEN" ) or os.environ.get( "hfjob" )
244 if hf_token:
245 login( token = hf_token)
246 training_args.hub_token = hf_token
247 logger.info( "Logged in to Hugging Face Hub" )
248 elif training_args.push_to_hub:
249 logger.warning( "HF_TOKEN not found in environment. Hub push will likely fail." )
250
251 trackio.init( project = training_args.output_dir, name = training_args.run_name)
252
253 logging.basicConfig(
254 format = " %(asctime)s - %(levelname)s - %(name)s - %(message)s " ,
255 datefmt = "%m/ %d /%Y %H:%M:%S" ,
256 handlers = [logging.StreamHandler(sys.stdout)],
257 )
258 if training_args.should_log:
259 transformers.utils.logging.set_verbosity_info()
260
261 log_level = training_args.get_process_log_level()
262 logger.setLevel(log_level)
263 transformers.utils.logging.set_verbosity(log_level)
264 transformers.utils.logging.enable_default_handler()
265 transformers.utils.logging.enable_explicit_format()
266
267 logger.info( f "Training/evaluation parameters { training_args } " )
268
269 # ---- Load dataset ----
270 dataset = load_dataset(
271 data_args.dataset_name,
272 data_args.dataset_config_name,
273 cache_dir = model_args.cache_dir,
274 trust_remote_code = model_args.trust_remote_code,
275 )
276
277 if "train" not in dataset:
278 if len (dataset.keys()) == 1 :
279 only_split = list (dataset.keys())[ 0 ]
280 dataset[only_split] = dataset[only_split].shuffle( seed = training_args.seed)
281 dataset = dataset[only_split].train_test_split( test_size = data_args.train_val_split or 0.1 )
282 dataset = { "train" : dataset[ "train" ], "validation" : dataset[ "test" ]}
283 else :
284 raise ValueError ( f "No 'train' split found. Available: { list (dataset.keys()) } " )
285 elif "validation" not in dataset and "test" not in dataset:
286 dataset[ "train" ] = dataset[ "train" ].shuffle( seed = training_args.seed)
287 split = dataset[ "train" ].train_test_split(
288 test_size = data_args.train_val_split or 0.1 , seed = training_args.seed
289 )
290 dataset[ "train" ] = split[ "train" ]
291 dataset[ "validation" ] = split[ "test" ]
292
293 if data_args.max_train_samples is not None :
294 n = min (data_args.max_train_samples, len (dataset[ "train" ]))
295 dataset[ "train" ] = dataset[ "train" ].select( range (n))
296 logger.info( f "Truncated training set to { n } samples" )
297 eval_key = "validation" if "validation" in dataset else "test"
298 if data_args.max_eval_samples is not None and eval_key in dataset:
299 n = min (data_args.max_eval_samples, len (dataset[eval_key]))
300 dataset[eval_key] = dataset[eval_key].select( range (n))
301 logger.info( f "Truncated eval set to { n } samples" )
302
303 # ---- Detect model family (SAM vs SAM2) and load processor/model ----
304 model_id = model_args.model_name_or_path.lower()
305 is_sam2 = "sam2" in model_id
306
307 if is_sam2:
308 from transformers import Sam2Processor, Sam2Model
309 processor = Sam2Processor.from_pretrained(model_args.model_name_or_path)
310 model = Sam2Model.from_pretrained(model_args.model_name_or_path)
311 else :
312 from transformers import SamProcessor, SamModel
313 processor = SamProcessor.from_pretrained(model_args.model_name_or_path)
314 model = SamModel.from_pretrained(model_args.model_name_or_path)
315
316 if model_args.freeze_vision_encoder:
317 for name, param in model.named_parameters():
318 if name.startswith( "vision_encoder" ):
319 param.requires_grad_( False )
320 if model_args.freeze_prompt_encoder:
321 for name, param in model.named_parameters():
322 if name.startswith( "prompt_encoder" ):
323 param.requires_grad_( False )
324
325 trainable = sum (p.numel() for p in model.parameters() if p.requires_grad)
326 total = sum (p.numel() for p in model.parameters())
327 logger.info( f "Trainable params: { trainable :,} / { total :,} ( { 100 * trainable / total :.1f} %)" )
328
329 # ---- Build datasets ----
330 prompt_col = data_args.prompt_column_name if data_args.prompt_column_name else None
331 ds_kwargs = dict (
332 processor = processor,
333 prompt_type = data_args.prompt_type,
334 image_col = data_args.image_column_name,
335 mask_col = data_args.mask_column_name,
336 prompt_col = prompt_col,
337 bbox_col = data_args.bbox_column_name,
338 point_col = data_args.point_column_name,
339 )
340
341 train_dataset = SAMSegmentationDataset( dataset = dataset[ "train" ], ** ds_kwargs)
342 eval_dataset = None
343 if eval_key in dataset:
344 eval_dataset = SAMSegmentationDataset( dataset = dataset[eval_key], ** ds_kwargs)
345
346 # ---- Train ----
347 trainer = Trainer(
348 model = model,
349 args = training_args,
350 train_dataset = train_dataset if training_args.do_train else None ,
351 eval_dataset = eval_dataset if training_args.do_eval else None ,
352 data_collator = collate_fn,
353 compute_loss_func = compute_loss,
354 )
355
356 if training_args.do_train:
357 train_result = trainer.train( resume_from_checkpoint = training_args.resume_from_checkpoint)
358 trainer.save_model()
359 trainer.log_metrics( "train" , train_result.metrics)
360 trainer.save_metrics( "train" , train_result.metrics)
361 trainer.save_state()
362
363 if training_args.do_eval and eval_dataset is not None :
364 metrics = trainer.evaluate()
365 trainer.log_metrics( "eval" , metrics)
366 trainer.save_metrics( "eval" , metrics)
367
368 trackio.finish()
369
370 kwargs = {
371 "finetuned_from" : model_args.model_name_or_path,
372 "dataset" : data_args.dataset_name,
373 "tags" : [ "image-segmentation" , "vision" , "sam" ],
374 }
375 if training_args.push_to_hub:
376 trainer.push_to_hub( ** kwargs)
377 else :
378 trainer.create_model_card( ** kwargs)
379
380
381 if __name__ == "__main__" :
382 main()