Views
No views yet
1import kagglehub
2import os
3os.environ["TOKENIZED_CACHE_DIR"] = kagglehub.dataset_download("gwendaltsang/tokenized-dataset/versions/3")Using Colab cache for faster access to the 'tokenized-dataset' dataset.1import os
2import time
3import warnings
4from pathlib import Path
5
6os.environ.setdefault("PJRT_DEVICE", "TPU")
7os.environ.setdefault("TF_CPP_MIN_LOG_LEVEL", "2")
8os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
9os.environ.setdefault("OMP_NUM_THREADS", "1")
10os.environ.setdefault("MKL_NUM_THREADS", "1")
11
12warnings.filterwarnings(
13 "ignore",
14 message=r"Transparent hugepages are not enabled\..*",
15 category=UserWarning,
16 module=r"jax\._src\.cloud_tpu_init",
17)
18
19SCRIPT_START = time.perf_counter()
20BACKUP_AFTER_SECONDS = 6590
21EPOCHS = 1
22
23import torch
24from datasets import load_from_disk
25from torch.utils.data import DataLoader
26from transformers import CamembertForMaskedLM, CamembertTokenizerFast
27
28import torch_xla
29import torch_xla.core.xla_model as xm
30import torch_xla.distributed.parallel_loader as pl
31from torch_xla.amp import syncfree
32
33
34def env_int(name: str, default: int) -> int:
35 return int(os.getenv(name, default))
36
37
38def env_float(name: str, default: float) -> float:
39 return float(os.getenv(name, default))
40
41
42def env_str(name: str, default: str) -> str:
43 return os.getenv(name, default)
44
45
46MODEL_NAME = env_str("MODEL_NAME", "camembert-base")
47TOKENIZED_CACHE_DIR = env_str("TOKENIZED_CACHE_DIR", "")
48SAVE_DIR = "/content/drive/MyDrive"
49
50BATCH_SIZE = env_int("BATCH_SIZE", 256)
51LEARNING_RATE = env_float("LEARNING_RATE", 5e-5)
52WEIGHT_DECAY = env_float("WEIGHT_DECAY", 0.01)
53WARMUP_RATIO = env_float("WARMUP_RATIO", 0.01)
54MAX_GRAD_NORM = env_float("MAX_GRAD_NORM", 1.0)
55MLM_PROBABILITY = env_float("MLM_PROBABILITY", 0.15)
56
57DATALOADER_WORKERS = env_int("DATALOADER_NUM_WORKERS", 16)
58PREFETCH_FACTOR = env_int("PREFETCH_FACTOR", 8)
59TORCH_NUM_THREADS = env_int("TORCH_NUM_THREADS", 4)
60
61XLA_LOADER_PREFETCH_SIZE = env_int("XLA_LOADER_PREFETCH_SIZE", 32)
62XLA_DEVICE_PREFETCH_SIZE = env_int("XLA_DEVICE_PREFETCH_SIZE", 16)
63XLA_TRANSFER_THREADS = env_int("XLA_TRANSFER_THREADS", 4)
64
65
66def make_optimizer(model: torch.nn.Module) -> torch.optim.Optimizer:
67 no_decay = ("bias", "LayerNorm.weight", "layer_norm.weight")
68
69 decay_params = [
70 p
71 for n, p in model.named_parameters()
72 if p.requires_grad and not any(nd in n for nd in no_decay)
73 ]
74 nodecay_params = [
75 p
76 for n, p in model.named_parameters()
77 if p.requires_grad and any(nd in n for nd in no_decay)
78 ]
79
80 return syncfree.AdamW(
81 [
82 {"params": decay_params, "weight_decay": WEIGHT_DECAY},
83 {"params": nodecay_params, "weight_decay": 0.0},
84 ],
85 lr=LEARNING_RATE,
86 betas=(0.9, 0.999),
87 eps=1e-6,
88 )
89
90
91def make_scheduler(
92 optimizer: torch.optim.Optimizer,
93 total_steps: int,
94) -> torch.optim.lr_scheduler.LambdaLR:
95 warmup_steps = int(total_steps * WARMUP_RATIO)
96
97 def lr_lambda(step: int) -> float:
98 warmup = step / max(1, warmup_steps)
99 decay = (total_steps - step) / max(1, total_steps - warmup_steps)
100 return min(warmup, max(0.0, decay))
101
102 return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
103
104
105def make_mlm_batch(
106 input_ids: torch.Tensor,
107 special_tokens_mask: torch.Tensor,
108 mask_token_id: int,
109 vocab_size: int,
110) -> tuple[torch.Tensor, torch.Tensor]:
111 rand_mask = torch.rand(input_ids.shape, device=input_ids.device)
112 masked = (rand_mask < MLM_PROBABILITY) & ~special_tokens_mask.bool()
113
114 labels = torch.where(masked, input_ids, torch.full_like(input_ids, -100))
115
116 replace_rand = torch.rand(input_ids.shape, device=input_ids.device)
117 replace_with_mask = masked & (replace_rand < 0.8)
118 replace_with_random = masked & (replace_rand >= 0.8) & (replace_rand < 0.9)
119
120 random_words = torch.randint(
121 low=0,
122 high=vocab_size,
123 size=input_ids.shape,
124 device=input_ids.device,
125 dtype=input_ids.dtype,
126 )
127
128 masked_input_ids = torch.where(
129 replace_with_mask,
130 torch.full_like(input_ids, mask_token_id),
131 input_ids,
132 )
133 masked_input_ids = torch.where(
134 replace_with_random,
135 random_words,
136 masked_input_ids,
137 )
138
139 return masked_input_ids, labels
140
141
142def save_model_weights_checkpoint(
143 model: torch.nn.Module,
144 tokenizer: CamembertTokenizerFast,
145 checkpoint_dir: str | Path,
146) -> None:
147 checkpoint_dir = Path(checkpoint_dir)
148 checkpoint_dir.mkdir(parents=True, exist_ok=True)
149
150 torch_xla.sync()
151
152 model_to_save = model.module if hasattr(model, "module") else model
153 cpu_state_dict = {
154 name: tensor.detach().cpu()
155 for name, tensor in model_to_save.state_dict().items()
156 }
157
158 model_to_save.save_pretrained(
159 str(checkpoint_dir),
160 state_dict=cpu_state_dict,
161 safe_serialization=True,
162 )
163 tokenizer.save_pretrained(str(checkpoint_dir))
164
165 del cpu_state_dict
166
167
168def train() -> None:
169 torch.set_num_threads(TORCH_NUM_THREADS)
170
171 device = torch_xla.device()
172
173 tokenizer = CamembertTokenizerFast.from_pretrained(MODEL_NAME)
174 mask_token_id = int(tokenizer.mask_token_id)
175 vocab_size = int(tokenizer.vocab_size)
176
177 dataset = load_from_disk(str(Path(TOKENIZED_CACHE_DIR)))
178 dataset.set_format(
179 type="torch",
180 columns=["input_ids", "attention_mask", "special_tokens_mask"],
181 )
182
183 train_loader = DataLoader(
184 dataset,
185 batch_size=BATCH_SIZE,
186 shuffle=True,
187 drop_last=True,
188 num_workers=DATALOADER_WORKERS,
189 persistent_workers=True,
190 prefetch_factor=PREFETCH_FACTOR,
191 )
192
193 model = CamembertForMaskedLM.from_pretrained(
194 MODEL_NAME,
195 use_safetensors=True,
196 ).to(device)
197 model.train()
198
199 optimizer = make_optimizer(model)
200
201 steps_per_epoch = len(train_loader)
202 scheduler = make_scheduler(optimizer, steps_per_epoch * EPOCHS)
203
204 xla_loader_kwargs = {
205 "loader_prefetch_size": XLA_LOADER_PREFETCH_SIZE,
206 "device_prefetch_size": XLA_DEVICE_PREFETCH_SIZE,
207 "host_to_device_transfer_threads": XLA_TRANSFER_THREADS,
208 }
209
210 optimizer.zero_grad(set_to_none=True)
211
212 for _ in range(EPOCHS):
213 device_loader = pl.MpDeviceLoader(
214 train_loader,
215 device,
216 **xla_loader_kwargs,
217 )
218
219 for step, batch in enumerate(device_loader, start=1):
220 input_ids, labels = make_mlm_batch(
221 input_ids=batch["input_ids"],
222 special_tokens_mask=batch["special_tokens_mask"],
223 mask_token_id=mask_token_id,
224 vocab_size=vocab_size,
225 )
226
227 with torch.autocast("xla", dtype=torch.bfloat16):
228 loss = model(
229 input_ids=input_ids,
230 attention_mask=batch["attention_mask"],
231 labels=labels,
232 ).loss
233
234 loss.backward()
235 torch.nn.utils.clip_grad_norm_(model.parameters(), MAX_GRAD_NORM)
236
237 xm.optimizer_step(optimizer)
238 scheduler.step()
239 optimizer.zero_grad(set_to_none=True)
240
241 if time.perf_counter() - SCRIPT_START >= BACKUP_AFTER_SECONDS:
242 save_model_weights_checkpoint(
243 model=model,
244 tokenizer=tokenizer,
245 checkpoint_dir=SAVE_DIR,
246 )
247 print(
248 f"[Backup 1h50] saved_dir={SAVE_DIR} "
249 f"step={step}/{steps_per_epoch} "
250 f"loss={float(loss.detach().cpu()):.6f}",
251 flush=True,
252 )
253 return
254
255 save_model_weights_checkpoint(
256 model=model,
257 tokenizer=tokenizer,
258 checkpoint_dir=SAVE_DIR,
259 )
260 print(f"[Final] saved_dir={SAVE_DIR} step={steps_per_epoch}/{steps_per_epoch}", flush=True)
261
262
263def main() -> None:
264 train()
265
266
267if __name__ == "__main__":
268 main()1Warning: You are sending unauthenticated requests to the HF Hub. Please set a HF_TOKEN to enable higher rate limits and faster downloads.
2WARNING:huggingface_hub.utils._http:Warning: You are sending unauthenticated requests to the HF Hub. Please set a HF_TOKEN to enable higher rate limits and faster downloads.
3tokenizer_config.json: 0%| | 0.00/25.0 [00:00<?, ?B/s]sentencepiece.bpe.model: 0%| | 0.00/811k [00:00<?, ?B/s]tokenizer.json: 0%| | 0.00/1.40M [00:00<?, ?B/s]config.json: 0%| | 0.00/508 [00:00<?, ?B/s]model.safetensors: 0%| | 0.00/445M [00:00<?, ?B/s]Loading weights: 0%| | 0/202 [00:00<?, ?it/s][transformers] CamembertForMaskedLM LOAD REPORT from: camembert-base
4Key | Status | |
5----------------------------+------------+--+-
6roberta.pooler.dense.bias | UNEXPECTED | |
7roberta.pooler.dense.weight | UNEXPECTED | |
8
9Notes:
10- UNEXPECTED: can be ignored when loading from different task/architecture; not ok if you expect identical arch.
11Writing model shards: 0%| | 0/1 [00:00<?, ?it/s][Backup 1h50] saved_dir=/content/drive/MyDrive step=13375/17516 loss=1.800222