From 61c36a0d988f4ca9393f9aba471acbda58696b8d Mon Sep 17 00:00:00 2001 From: Yifeng He Date: Sat, 7 Feb 2026 19:25:45 +0000 Subject: [PATCH 1/5] fix: clone input_ids to prevent in-place modification error in DDP training --- modeling/model.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/modeling/model.py b/modeling/model.py index 97ab6e3..62b0bd6 100644 --- a/modeling/model.py +++ b/modeling/model.py @@ -226,12 +226,14 @@ def compute_loss( if self.contra_mode == ContraMode.GROUPED: return self._compute_grouped_loss(model, inputs, return_outputs) - # Move inputs to the device the model lives on (supports DDP) + # Move inputs to the device the model lives on (supports DDP). + # Clone input_ids tensors to avoid in-place modification conflicts + # during backward when the MLM head or embeddings share storage. device = model.device - code_input_ids = inputs["code_input_ids"].to(device) + code_input_ids = inputs["code_input_ids"].to(device).clone() code_attention_mask = inputs["code_attention_mask"].to(device) code_labels = inputs["code_labels"].to(device) - aug_input_ids = inputs["aug_input_ids"].to(device) + aug_input_ids = inputs["aug_input_ids"].to(device).clone() aug_attention_mask = inputs["aug_attention_mask"].to(device) aug_labels = inputs["aug_labels"].to(device) @@ -295,10 +297,10 @@ def _compute_grouped_loss(self, model, inputs, return_outputs=False): - group_sizes: [B] """ device = model.device - code_input_ids = inputs["code_input_ids"].to(device) + code_input_ids = inputs["code_input_ids"].to(device).clone() code_attention_mask = inputs["code_attention_mask"].to(device) code_labels = inputs["code_labels"].to(device) - aug_input_ids = inputs["aug_input_ids"].to(device) + aug_input_ids = inputs["aug_input_ids"].to(device).clone() aug_attention_mask = inputs["aug_attention_mask"].to(device) aug_labels = inputs["aug_labels"].to(device) group_sizes = inputs["group_sizes"].to(device) From 29a98d5aca532d0adb78f1f8562c416586ffab1a Mon Sep 17 00:00:00 2001 From: Yifeng He Date: Sat, 7 Feb 2026 19:36:00 +0000 Subject: [PATCH 2/5] perf: enable HF datasets caching for tokenization map Replace lambdas with functools.partial so HF datasets can produce stable fingerprints for .map() calls. Move .shuffle() after .map() so the tokenization cache is reused across runs regardless of seed. --- modeling/dataloader.py | 2 +- modeling/pretrain.py | 24 +++++++++++++----------- tests/test_grouped.py | 10 +++++----- 3 files changed, 19 insertions(+), 17 deletions(-) diff --git a/modeling/dataloader.py b/modeling/dataloader.py index 70d2ca7..5ddf034 100644 --- a/modeling/dataloader.py +++ b/modeling/dataloader.py @@ -66,7 +66,7 @@ def contra_data_collator(mlm_collator, features): return batch -def grouped_contra_data_collator(mlm_collator, features, max_num_augs): +def grouped_contra_data_collator(mlm_collator, max_num_augs, features): """Collate grouped samples where each item has 1 anchor + variable-count augmentations. Each feature dict contains: diff --git a/modeling/pretrain.py b/modeling/pretrain.py index a3c5ead..7679ad8 100644 --- a/modeling/pretrain.py +++ b/modeling/pretrain.py @@ -3,6 +3,7 @@ import hashlib import os from collections import defaultdict +from functools import partial import torch import torch.distributed as dist @@ -270,25 +271,26 @@ def main( if contra_mode == "grouped": # Regroup flat rows by function_id into {code, [aug_1, ..., aug_K]} grouped_dataset = regroup_dataset(dataset, max_num_augs=max_num_augs) - tokenized_datasets = grouped_dataset.shuffle(seed=seed).map( - lambda example: tokenize_grouped( - tokenizer, example, max_seq_length, max_num_augs + tokenized_datasets = grouped_dataset.map( + partial( + tokenize_grouped, + tokenizer, + max_seq_length=max_seq_length, + max_num_augs=max_num_augs, ), batched=True, num_proc=num_proc, - ) + ).shuffle(seed=seed) - collator_fn = lambda features: grouped_contra_data_collator( - mlm_collator, features, max_num_augs - ) + collator_fn = partial(grouped_contra_data_collator, mlm_collator, max_num_augs) else: - tokenized_datasets = dataset.shuffle(seed=seed).map( - lambda example: tokenize(tokenizer, example, max_seq_length=max_seq_length), + tokenized_datasets = dataset.map( + partial(tokenize, tokenizer, max_seq_length=max_seq_length), batched=True, num_proc=num_proc, - ) + ).shuffle(seed=seed) - collator_fn = lambda features: contra_data_collator(mlm_collator, features) + collator_fn = partial(contra_data_collator, mlm_collator) split_dataset = tokenized_datasets.train_test_split(test_size=0.1) train_dataset = split_dataset["train"] diff --git a/tests/test_grouped.py b/tests/test_grouped.py index b77e0da..37c1fb1 100644 --- a/tests/test_grouped.py +++ b/tests/test_grouped.py @@ -162,7 +162,7 @@ def test_output_shapes(self): mlm_collator = self._get_mlm_collator() seq_len = 16 features = _make_grouped_feature(seq_len, aug_counts=[2, 3]) - batch = grouped_contra_data_collator(mlm_collator, features, max_num_augs=6) + batch = grouped_contra_data_collator(mlm_collator, 6, features) B = 2 max_K = 3 # max(2, 3) @@ -177,7 +177,7 @@ def test_padding_has_no_mlm_labels(self): seq_len = 16 # Group 0: 1 aug, Group 1: 3 augs → max_K=3, group 0 has 2 padding slots features = _make_grouped_feature(seq_len, aug_counts=[1, 3]) - batch = grouped_contra_data_collator(mlm_collator, features, max_num_augs=6) + batch = grouped_contra_data_collator(mlm_collator, 6, features) # Group 0's padding slots are indices 1 and 2 in the flattened aug batch # (group 0 occupies slots 0..2, real=1, padding=slots 1,2) @@ -190,14 +190,14 @@ def test_group_sizes_correct(self): """group_sizes should reflect actual aug counts.""" mlm_collator = self._get_mlm_collator() features = _make_grouped_feature(16, aug_counts=[1, 2, 4]) - batch = grouped_contra_data_collator(mlm_collator, features, max_num_augs=6) + batch = grouped_contra_data_collator(mlm_collator, 6, features) assert batch["group_sizes"].tolist() == [1, 2, 4] def test_max_num_augs_truncation(self): """Features with more augs than max_num_augs get truncated.""" mlm_collator = self._get_mlm_collator() features = _make_grouped_feature(16, aug_counts=[5, 3]) - batch = grouped_contra_data_collator(mlm_collator, features, max_num_augs=2) + batch = grouped_contra_data_collator(mlm_collator, 2, features) B = 2 max_K = 2 @@ -208,7 +208,7 @@ def test_function_id_passed_through(self): """function_id should be present in the batch.""" mlm_collator = self._get_mlm_collator() features = _make_grouped_feature(16, aug_counts=[2, 1]) - batch = grouped_contra_data_collator(mlm_collator, features, max_num_augs=6) + batch = grouped_contra_data_collator(mlm_collator, 6, features) assert "function_id" in batch assert batch["function_id"].shape == (2,) From 65b98c5caccc577b18cccf960a16e824f48a12c8 Mon Sep 17 00:00:00 2001 From: Yifeng He Date: Sat, 7 Feb 2026 20:02:42 +0000 Subject: [PATCH 3/5] fix: merge dual forward passes into one to fix DDP in-place buffer error --- modeling/model.py | 88 +++++++++++++++++++++-------------------------- 1 file changed, 39 insertions(+), 49 deletions(-) diff --git a/modeling/model.py b/modeling/model.py index 62b0bd6..e4aaa5f 100644 --- a/modeling/model.py +++ b/modeling/model.py @@ -226,45 +226,36 @@ def compute_loss( if self.contra_mode == ContraMode.GROUPED: return self._compute_grouped_loss(model, inputs, return_outputs) - # Move inputs to the device the model lives on (supports DDP). - # Clone input_ids tensors to avoid in-place modification conflicts - # during backward when the MLM head or embeddings share storage. + # Concatenate code and aug inputs into a single batch so that DDP + # sees exactly one forward pass per backward (two separate forwards + # through DDP cause in-place version errors on internal buffers). device = model.device - code_input_ids = inputs["code_input_ids"].to(device).clone() + code_input_ids = inputs["code_input_ids"].to(device) code_attention_mask = inputs["code_attention_mask"].to(device) code_labels = inputs["code_labels"].to(device) - aug_input_ids = inputs["aug_input_ids"].to(device).clone() + aug_input_ids = inputs["aug_input_ids"].to(device) aug_attention_mask = inputs["aug_attention_mask"].to(device) aug_labels = inputs["aug_labels"].to(device) - # Forward pass for MLM - # use bi-encoder training, encode code and augmentation separately using self.model - code_outputs = model( - input_ids=code_input_ids, - attention_mask=code_attention_mask, - labels=code_labels, - output_hidden_states=True, - return_dict=True, - ) - code_hidden_states = code_outputs.hidden_states[-1] - code_embeddings = code_hidden_states[:, 0, :] + B = code_input_ids.size(0) + + all_input_ids = torch.cat([code_input_ids, aug_input_ids], dim=0) + all_attention_mask = torch.cat([code_attention_mask, aug_attention_mask], dim=0) + all_labels = torch.cat([code_labels, aug_labels], dim=0) - aug_outputs = model( - input_ids=aug_input_ids, - attention_mask=aug_attention_mask, - labels=aug_labels, + outputs = model( + input_ids=all_input_ids, + attention_mask=all_attention_mask, + labels=all_labels, output_hidden_states=True, return_dict=True, ) - aug_hidden_states = aug_outputs.hidden_states[-1] - aug_embeddings = aug_hidden_states[:, 0, :] - # Average MLM losses so the combined MLM term is on the same scale - # as the single contrastive term (~3-8 each), letting alpha express - # a genuine preference rather than compensating for a 2x scale artifact. - code_mlm_loss = code_outputs.loss - aug_mlm_loss = aug_outputs.loss - mlm_loss = (code_mlm_loss + aug_mlm_loss) / 2 + hidden_states = outputs.hidden_states[-1] + code_embeddings = hidden_states[:B, 0, :] + aug_embeddings = hidden_states[B:, 0, :] + + mlm_loss = outputs.loss # Compute contrastive loss between code and its augmentation if self.contra_mode == ContraMode.SUPCON: @@ -281,10 +272,9 @@ def compute_loss( self.temperature, ) - # Total loss with weighting (adjust alpha as needed) total_loss = mlm_loss + self.alpha * contrastive_loss - return (total_loss, code_outputs) if return_outputs else total_loss + return (total_loss, outputs) if return_outputs else total_loss def _compute_grouped_loss(self, model, inputs, return_outputs=False): """Compute loss for grouped multi-key contrast mode. @@ -297,36 +287,36 @@ def _compute_grouped_loss(self, model, inputs, return_outputs=False): - group_sizes: [B] """ device = model.device - code_input_ids = inputs["code_input_ids"].to(device).clone() + code_input_ids = inputs["code_input_ids"].to(device) code_attention_mask = inputs["code_attention_mask"].to(device) code_labels = inputs["code_labels"].to(device) - aug_input_ids = inputs["aug_input_ids"].to(device).clone() + aug_input_ids = inputs["aug_input_ids"].to(device) aug_attention_mask = inputs["aug_attention_mask"].to(device) aug_labels = inputs["aug_labels"].to(device) group_sizes = inputs["group_sizes"].to(device) - # Forward anchor - code_outputs = model( - input_ids=code_input_ids, - attention_mask=code_attention_mask, - labels=code_labels, - output_hidden_states=True, - return_dict=True, - ) - code_embeddings = code_outputs.hidden_states[-1][:, 0, :] # [B, D] + B = code_input_ids.size(0) + + # Single forward pass: concatenate anchors and augmentations to avoid + # DDP in-place buffer errors from two separate forward calls. + all_input_ids = torch.cat([code_input_ids, aug_input_ids], dim=0) + all_attention_mask = torch.cat([code_attention_mask, aug_attention_mask], dim=0) + all_labels = torch.cat([code_labels, aug_labels], dim=0) - # Forward all augmentations (flattened [B*max_K, seq_len]) - aug_outputs = model( - input_ids=aug_input_ids, - attention_mask=aug_attention_mask, - labels=aug_labels, + outputs = model( + input_ids=all_input_ids, + attention_mask=all_attention_mask, + labels=all_labels, output_hidden_states=True, return_dict=True, ) - aug_embeddings = aug_outputs.hidden_states[-1][:, 0, :] # [B*max_K, D] + + hidden_states = outputs.hidden_states[-1] + code_embeddings = hidden_states[:B, 0, :] # [B, D] + aug_embeddings = hidden_states[B:, 0, :] # [B*max_K, D] # MLM loss (padding augs have labels=-100, contribute 0) - mlm_loss = (code_outputs.loss + aug_outputs.loss) / 2 + mlm_loss = outputs.loss # Contrastive loss contrastive_loss = grouped_contrastive_loss( @@ -335,7 +325,7 @@ def _compute_grouped_loss(self, model, inputs, return_outputs=False): total_loss = mlm_loss + self.alpha * contrastive_loss - return (total_loss, code_outputs) if return_outputs else total_loss + return (total_loss, outputs) if return_outputs else total_loss def prediction_step(self, model, inputs, prediction_loss_only, ignore_keys=None): """ From fedc8038d1f2d6a40fac594c6a20831c4110601c Mon Sep 17 00:00:00 2001 From: Yifeng He Date: Sun, 8 Feb 2026 04:40:23 +0800 Subject: [PATCH 4/5] remove cap on 80 process --- modeling/common.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/modeling/common.py b/modeling/common.py index a42665f..ed9cbe3 100644 --- a/modeling/common.py +++ b/modeling/common.py @@ -9,7 +9,7 @@ def default_num_proc() -> int: """Return the default number of parallel workers, capped at available CPUs.""" - return min(os.cpu_count() or 1, MAX_NUM_PROC) + return os.cpu_count() or 1 def set_seed(seed): From 2e60959fd24690805d30e8ef1a5411f48c09692a Mon Sep 17 00:00:00 2001 From: Yifeng He Date: Sun, 8 Feb 2026 05:16:02 +0800 Subject: [PATCH 5/5] adjust batch size --- experiments/grouped/codebert.yaml | 4 ++-- experiments/grouped/contrabert_c.yaml | 4 ++-- experiments/grouped/contrabert_g.yaml | 4 ++-- experiments/grouped/graphcodebert.yaml | 4 ++-- 4 files changed, 8 insertions(+), 8 deletions(-) diff --git a/experiments/grouped/codebert.yaml b/experiments/grouped/codebert.yaml index addef07..8bcf1e9 100644 --- a/experiments/grouped/codebert.yaml +++ b/experiments/grouped/codebert.yaml @@ -5,9 +5,9 @@ dataset_path: "data/aug_csn.jsonl" model_name: "microsoft/codebert-base" -batch_size: 16 +batch_size: 64 num_epochs: 3 -gradient_accumulation_steps: 16 +gradient_accumulation_steps: 4 learning_rate: 2.0e-5 seed: 0 diff --git a/experiments/grouped/contrabert_c.yaml b/experiments/grouped/contrabert_c.yaml index 1f8e87c..d5f8f4d 100644 --- a/experiments/grouped/contrabert_c.yaml +++ b/experiments/grouped/contrabert_c.yaml @@ -6,9 +6,9 @@ dataset_path: "data/aug_csn.jsonl" model_name: "./saved_models/ContraBERT_C" tokenizer_name: "microsoft/codebert-base" -batch_size: 16 +batch_size: 64 num_epochs: 3 -gradient_accumulation_steps: 16 +gradient_accumulation_steps: 4 learning_rate: 2.0e-5 seed: 0 diff --git a/experiments/grouped/contrabert_g.yaml b/experiments/grouped/contrabert_g.yaml index 0b8db32..392420a 100644 --- a/experiments/grouped/contrabert_g.yaml +++ b/experiments/grouped/contrabert_g.yaml @@ -6,9 +6,9 @@ dataset_path: "data/aug_csn.jsonl" model_name: "./saved_models/ContraBERT_G" tokenizer_name: "microsoft/graphcodebert-base" -batch_size: 16 +batch_size: 64 num_epochs: 3 -gradient_accumulation_steps: 16 +gradient_accumulation_steps: 4 learning_rate: 2.0e-5 seed: 0 diff --git a/experiments/grouped/graphcodebert.yaml b/experiments/grouped/graphcodebert.yaml index 0e2d204..7e235df 100644 --- a/experiments/grouped/graphcodebert.yaml +++ b/experiments/grouped/graphcodebert.yaml @@ -5,9 +5,9 @@ dataset_path: "data/aug_csn.jsonl" model_name: "microsoft/graphcodebert-base" -batch_size: 16 +batch_size: 64 num_epochs: 3 -gradient_accumulation_steps: 16 +gradient_accumulation_steps: 4 learning_rate: 2.0e-5 seed: 0