(device_id)
| 23 | |
| 24 | |
| 25 | def train_worker(device_id): |
| 26 | model_id = 'damo/multi-modal_team-vit-large-patch14_multi-modal-similarity' |
| 27 | ckpt_dir = './ckpt' |
| 28 | os.makedirs(ckpt_dir, exist_ok=True) |
| 29 | # Use epoch=1 for faster training here |
| 30 | cfg = Config({ |
| 31 | 'framework': 'pytorch', |
| 32 | 'task': 'multi-modal-similarity', |
| 33 | 'pipeline': { |
| 34 | 'type': 'multi-modal-similarity' |
| 35 | }, |
| 36 | 'model': { |
| 37 | 'type': 'team-multi-modal-similarity' |
| 38 | }, |
| 39 | 'dataset': { |
| 40 | 'name': 'Caltech101', |
| 41 | 'class_num': 101 |
| 42 | }, |
| 43 | 'preprocessor': {}, |
| 44 | 'train': { |
| 45 | 'epoch': 1, |
| 46 | 'batch_size': 32, |
| 47 | 'ckpt_dir': ckpt_dir |
| 48 | }, |
| 49 | 'evaluation': { |
| 50 | 'batch_size': 64 |
| 51 | } |
| 52 | }) |
| 53 | cfg_file = '{}/{}'.format(ckpt_dir, ModelFile.CONFIGURATION) |
| 54 | cfg.dump(cfg_file) |
| 55 | |
| 56 | train_dataset = MsDataset.load( |
| 57 | cfg.dataset.name, |
| 58 | namespace='modelscope', |
| 59 | split='train', |
| 60 | download_mode=DownloadMode.FORCE_REDOWNLOAD).to_hf_dataset() |
| 61 | train_dataset = train_dataset.with_transform(train_mapping) |
| 62 | val_dataset = MsDataset.load( |
| 63 | cfg.dataset.name, |
| 64 | namespace='modelscope', |
| 65 | split='validation', |
| 66 | download_mode=DownloadMode.FORCE_REDOWNLOAD).to_hf_dataset() |
| 67 | val_dataset = val_dataset.with_transform(val_mapping) |
| 68 | |
| 69 | default_args = dict( |
| 70 | cfg_file=cfg_file, |
| 71 | model=model_id, |
| 72 | device_id=device_id, |
| 73 | data_collator=collate_fn, |
| 74 | train_dataset=train_dataset, |
| 75 | val_dataset=val_dataset) |
| 76 | |
| 77 | trainer = build_trainer( |
| 78 | name=Trainers.image_classification_team, default_args=default_args) |
| 79 | trainer.train() |
| 80 | trainer.evaluate() |
| 81 | |
| 82 |
no test coverage detected
searching dependent graphs…