Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions examples/qwen2.5-7B-distill_megatron/distill_megatron.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ checkpoint_config:
save_steps: 100
logging_steps: 1
resume_from_checkpoint: false
eval_steps: 50

student_pretrain: Qwen/Qwen2.5-7B-Instruct
teacher_pretrain: Qwen/Qwen2.5-14B-Instruct
Expand All @@ -34,6 +35,12 @@ max_grad_norm: 1.0
question_key: question_zh
answer_key: answer_zh

validation:
data_args:
file_name:
- data/GSM8K_zh.json
template: qwen2_5

student:
model_args:
attn_implementation: fa2
Expand Down Expand Up @@ -64,6 +71,7 @@ student:
recompute_granularity: full
use_sequence_packing: True
device_mapping: list(range(0,8))
infer_batch_size: 2

teacher:
model_args:
Expand Down
4 changes: 4 additions & 0 deletions roll/pipeline/distill/distill_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,10 @@ class DistillConfig(BaseConfig):
default_factory=WorkerConfig,
metadata={"help": "Configuration for the teacher's role."}
)
validation: WorkerConfig = field(
default=None,
metadata={"help": "Configuration for the validation."}
)

# data related
question_key: str = field(
Expand Down
45 changes: 45 additions & 0 deletions roll/pipeline/distill/distill_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@

import datasets
import ray
import numpy as np
import torch
from torch.utils.data import DataLoader
from codetiming import Timer
Expand Down Expand Up @@ -177,6 +178,14 @@ def __init__(self, pipeline_config: DistillConfig):
raise ValueError("No dataset paths provided")
print(f'load_dataset_paths: {chr(10)} {chr(10).join(dataset_paths)}')
dataset = datasets.load_dataset('json', data_files=dataset_paths)['train']

val_dataset = None
if self.pipeline_config.validation and self.pipeline_config.validation.data_args:
val_dataset_paths = self.pipeline_config.validation.data_args.file_name
if not val_dataset_paths:
raise ValueError("No val dataset paths provided")
print(f'load_dataset_paths: {chr(10)} {chr(10).join(val_dataset_paths)}')
val_dataset = datasets.load_dataset("json", data_files=val_dataset_paths)["train"]

# Currently, only models where the student and teacher are of the same type are supported.
self.tokenizer = default_tokenizer_provider(model_args=self.pipeline_config.student.model_args)
Expand Down Expand Up @@ -231,6 +240,24 @@ def __init__(self, pipeline_config: DistillConfig):
data_collator,
num_proc=self.pipeline_config.student.training_args.dataloader_num_workers)

if val_dataset:
val_dataset = preprocess_dataset(
val_dataset,
self.tokenizer,
pipeline_config
)

self.val_dataloader = DataLoader(
dataset=val_dataset,
batch_size=self.pipeline_config.student.infer_batch_size *\
self.pipeline_config.student.training_args.gradient_accumulation_steps *\
self.student.get_rank_info(0).dp_size,
shuffle=False,
drop_last=True,
num_workers=self.pipeline_config.student.training_args.dataloader_num_workers,
collate_fn=data_collator
)

self.set_checkpoint_clusters(self.student)

@torch.no_grad()
Expand All @@ -248,6 +275,12 @@ def run(self):
logger.info(f"pipeline step {global_step} start...")

metrics_mgr.clear_metrics()

if self.val_dataloader and global_step % self.pipeline_config.eval_steps == 0:
with Timer(name="val") as val_timer:
val_metrics = self.val()
metrics_mgr.add_reduced_metrics(val_metrics)
metrics_mgr.add_metric("time/val", val_timer.last)

batch: DataProto = DataProto.from_single_dict(batch_dict)
batch.meta_info = {"global_step": global_step, "is_offload_states": False, "is_offload_optimizer_states_in_train_step": False}
Expand Down Expand Up @@ -293,3 +326,15 @@ def run(self):
logger.info(f"pipeline step {global_step} finished")
global_step += 1
logger.info("pipeline complete!")

@torch.no_grad()
def val(self):
val_loss_list = []
for batch_dict in self.val_dataloader:
batch: DataProto = DataProto.from_single_dict(batch_dict)
batch.meta_info = {"is_offload_optimizer_states_in_train_step": False}
val_metrics_refs = self.student.val_step(batch, blocking=False)
val_metrics = DataProto.materialize_concat(data_refs=val_metrics_refs)
val_metrics = val_metrics.meta_info.pop("metrics", {})
val_loss_list.append(val_metrics[f"student/val_loss"])
return {"student/val_loss": np.concatenate(val_loss_list)}
20 changes: 20 additions & 0 deletions roll/pipeline/distill/distill_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -167,6 +167,26 @@ def loss_func(self, data: DataProto, output_tensor: torch.Tensor):
}
return loss, student_metrics

@register(Dispatch.DP_MP_DISPATCH_FIRST, clear_cache=False)
def val_step(self, data: DataProto):
data = data.to(current_platform.device_type)
data.meta_info["micro_batch_size"] = self.worker_config.infer_batch_size
data = self.strategy.get_data_input(data)
if "labels" in data.batch.keys():
# rename key: labels -> labels_for_loss
data.batch.rename_key_("labels", "labels_for_loss")
metrics = self.strategy.forward_step(batch=data, forward_func=self.loss_func_for_eval)
output = DataProto(meta_info={"metrics": metrics}).to("cpu")
return output

def loss_func_for_eval(self, data: DataProto, output_tensor: torch.Tensor):
labels = data.batch['labels_for_loss']
gpt_loss, _ = self.strategy.op_compute_language_loss_from_logits(output_tensor, labels)
student_metrics = {
"student/val_loss": gpt_loss.detach().item(),
}
return gpt_loss, student_metrics

@register(dispatch_mode=Dispatch.ONE_TO_ALL)
def do_checkpoint(self, global_step):
with Timer("do_checkpoint") as total_timer:
Expand Down