Download scripts/train.py from OneScience-Group/RemoteCLIP: direct link, hf CLI and curl.
- Browser
- Download file 4.68 kB
-
https://huggingface.co/OneScience-Group/RemoteCLIP/resolve/7002f4e4517867ddba3bfa7735db916bd0ff155f/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/RemoteCLIP@7002f4e4517867ddba3bfa7735db916bd0ff155f/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/RemoteCLIP/resolve/7002f4e4517867ddba3bfa7735db916bd0ff155f/scripts/train.py
4.68 kB
| """Train compact RemoteCLIP with bidirectional InfoNCE.""" | |
| import importlib.util | |
| import json | |
| import os | |
| import random | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import yaml | |
| from torch import distributed as dist | |
| from torch.nn.parallel import DistributedDataParallel | |
| from torch.utils.data import DataLoader, Dataset, DistributedSampler | |
| ROOT = Path(__file__).resolve().parents[1] | |
| def load_model_class(): | |
| spec = importlib.util.spec_from_file_location("remoteclip_model", ROOT / "model" / "remoteclip.py") | |
| module = importlib.util.module_from_spec(spec) | |
| spec.loader.exec_module(module) | |
| return module.RemoteCLIP | |
| class PairDataset(Dataset): | |
| def __init__(self, path): | |
| archive = np.load(path) | |
| self.images = archive["train_images"] | |
| self.tokens = archive["train_tokens"] | |
| self.data_source = str(archive["data_source"]) if "data_source" in archive.files else "unknown" | |
| self.protocol = str(archive["protocol"]) if "protocol" in archive.files else "unknown" | |
| def __len__(self): | |
| return len(self.images) | |
| def __getitem__(self, index): | |
| return torch.from_numpy(self.images[index]), torch.from_numpy(self.tokens[index]) | |
| def main(): | |
| with (ROOT / "conf" / "config.yaml").open(encoding="utf-8") as handle: | |
| config = yaml.safe_load(handle) | |
| data_path = ROOT / config["data"]["path"] | |
| if not data_path.exists(): | |
| raise FileNotFoundError( | |
| f"Missing training data: {data_path.relative_to(ROOT)}. " | |
| "Run `python scripts/fake_data.py` first." | |
| ) | |
| world_size = int(os.environ.get("WORLD_SIZE", "1")) | |
| local_rank = int(os.environ.get("LOCAL_RANK", "0")) | |
| if world_size > 1: | |
| dist.init_process_group("nccl" if torch.cuda.is_available() else "gloo") | |
| device = torch.device(f"cuda:{local_rank}" if torch.cuda.is_available() else "cpu") | |
| if device.type == "cuda": | |
| torch.cuda.set_device(local_rank) | |
| seed = config["seed"] + local_rank | |
| random.seed(seed) | |
| np.random.seed(seed) | |
| torch.manual_seed(seed) | |
| dataset = PairDataset(data_path) | |
| sampler = DistributedSampler(dataset, shuffle=True) if world_size > 1 else None | |
| loader = DataLoader( | |
| dataset, | |
| batch_size=config["train"]["batch_size"], | |
| shuffle=sampler is None, | |
| sampler=sampler, | |
| num_workers=config["train"]["num_workers"], | |
| ) | |
| RemoteCLIP = load_model_class() | |
| model = RemoteCLIP( | |
| vocabulary_size=config["data"]["vocabulary_size"], | |
| context_length=config["data"]["context_length"], | |
| **config["model"], | |
| ).to(device) | |
| if world_size > 1: | |
| model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) | |
| optimizer = torch.optim.AdamW( | |
| model.parameters(), | |
| lr=config["train"]["learning_rate"], | |
| weight_decay=config["train"]["weight_decay"], | |
| ) | |
| history = [] | |
| for epoch in range(config["train"]["epochs"]): | |
| if sampler is not None: | |
| sampler.set_epoch(epoch) | |
| model.train() | |
| total = 0.0 | |
| for images, tokens in loader: | |
| output = model(images.to(device), tokens.to(device)) | |
| optimizer.zero_grad(set_to_none=True) | |
| output["loss"].backward() | |
| optimizer.step() | |
| total += output["loss"].item() | |
| loss = total / len(loader) | |
| history.append({"epoch": epoch + 1, "contrastive_loss": loss}) | |
| if local_rank == 0: | |
| print(f"epoch={epoch + 1} contrastive_loss={loss:.6f}") | |
| if local_rank == 0: | |
| checkpoint = ROOT / config["paths"]["checkpoint"] | |
| metrics = ROOT / config["paths"]["training_metrics"] | |
| checkpoint.parent.mkdir(parents=True, exist_ok=True) | |
| metrics.parent.mkdir(parents=True, exist_ok=True) | |
| base_model = model.module if hasattr(model, "module") else model | |
| torch.save( | |
| { | |
| "model": base_model.state_dict(), | |
| "config": config, | |
| "data_source": dataset.data_source, | |
| "protocol": dataset.protocol, | |
| }, | |
| checkpoint, | |
| ) | |
| metrics.write_text( | |
| json.dumps( | |
| {"history": history, "data_source": dataset.data_source, "protocol": dataset.protocol}, | |
| indent=2, | |
| ) | |
| + "\n", | |
| encoding="utf-8", | |
| ) | |
| print( | |
| f"checkpoint={checkpoint.relative_to(ROOT)} data_source={dataset.data_source} " | |
| f"protocol={dataset.protocol}" | |
| ) | |
| if world_size > 1: | |
| dist.destroy_process_group() | |
| if __name__ == "__main__": | |
| main() | |