deepspeedai / deepspeedai/DeepSpeed
[BUG]
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 43.1k
- Forks
- 5k
- Avg merge
- 4d 15h
- Merged PRs (30d)
- 112
Description
Describe the bug
torchrun & deepspeed stage 2: memory imbalance(Using SFTtrainer)
If I put the evaluation mode, there is this much difference, and if I don't put the evaluation mode, the usage difference between gpu 0 and gpu 1 is about 10gb.
I think the problem is that the distributed sampler is not working properly when working with the sfttrainer in deepspeed, or the cache is constantly piling up.
| 0 NVIDIA A100-SXM4-80GB On | 00000000:07:00.0 Off | 0 |
| N/A 59C P0 372W / 400W | 43591MiB / 81920MiB | 95% Default |
| | | Disabled |
+-----------------------------------------+----------------------+----------------------+
| 1 NVIDIA A100-SXM4-80GB On | 00000000:0A:00.0 Off | 0 |
| N/A 41C P0 123W / 400W | 77117MiB / 81920MiB | 99% Default |
Source Code
import dataclasses
import wandb
import json
import logging
import os
import warnings
import psutil
import gc
from typing import Dict, Optional, List, Union, Any
from functools import partial
import torch
from torch.utils.data import Dataset, DataLoader
import transformers
from trl import SFTTrainer, SFTConfig
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
DataCollatorForLanguageModeling,
HfArgumentParser,
set_seed,
get_scheduler
)
from peft import LoraConfig, get_peft_model, TaskType, prepare_model_for_kbit_training
import deepspeed
from deepspeed import zero
from deepspeed.runtime.zero.partition_parameters import ZeroParamStatus
import datetime
import torch.distributed as dist
import time
import numpy as np
import torch.nn.functional as F
from torch.utils.data.distributed import DistributedSampler
warnings.filterwarnings("ignore", message="Detected kernel version.below the recommended minimum.")
hf_token = os.getenv("HF_TOKEN")
logger = logging.getLogger(name)
logging.basicConfig(
format="%(asctime)s %(levelname)s [%(name)s] %(message)s",
level=logging.INFO,
datefmt="%Y-%m-%d %H:%M:%S"
)
INTRO_BLURB = "Below is an instruction that describes a task. Write a response that appropriately completes the request."
INSTRUCTION_KEY = "### Instruction:"
INPUT_KEY = "Input:"
RESPONSE_KEY = "### Response:"
END_KEY = "### End"
RESPONSE_KEY_NL = f"{RESPONSE_KEY}\n"
PROMPT_NO_INPUT_FORMAT = """{intro}
{instruction_key}
{instruction}
{response_key}
{response}
{end_key}""".format(
intro=INTRO_BLURB,
instruction_key=INSTRUCTION_KEY,
instruction="{instruction}",
response_key=RESPONSE_KEY,
response="{response}",
end_key=END_KEY,
)
PROMPT_WITH_INPUT_FORMAT = """{intro}
{instruction_key}
{instruction}
{input_key}
{input}
{response_key}
{response}
{end_key}""".format(
intro=INTRO_BLURB,
instruction_key=INSTRUCTION_KEY,
instruction="{instruction}",
input_key=INPUT_KEY,
input="{input}",
response_key=RESPONSE_KEY,
response="{response}",
end_key=END_KEY,
)
@dataclasses.dataclass
class ModelArguments:
model_name_or_path: Optional[str] = dataclasses.field(default="Qwen/Qwen2.5-7B")
@dataclasses.dataclass
class DataArguments:
data_path: str = dataclasses.field(default="databricks/databricks-dolly-15k", metadata={"help": "Path to the training data."})
eval_data_path: str = dataclasses.field(default=None, metadata={"help": "Path to the evaluation data."})
lazy_preprocess: bool = False
test_size: int = dataclasses.field(default=1000, metadata={"help": "Number of examples for test set"})
max_samples: Optional[int] = dataclasses.field(default=None, metadata={"help": "Max number of samples to use from the dataset"})
@dataclasses.dataclass
class CustomTrainingArguments(SFTConfig):
optim: str = dataclasses.field(default="adamw_torch")
use_lora: bool = dataclasses.field(default=True, metadata={"help": "Whether to use LoRA"})
output_dir: str = dataclasses.field(default="./output", metadata={"help": "Output directory"})
packing: bool = dataclasses.field(default=False, metadata={"help": "Whether to pack sequences"})
max_seq_length: Optional[int] = dataclasses.field(default=2048, metadata={"help": "Maximum sequence length"})
ddp_find_unused_parameters: bool = dataclasses.field(default=False, metadata={"help": "Whether to find unused parameters in DDP"})
per_device_train_batch_size: int = dataclasses.field(default=4, metadata={"help": "Batch size per device for training"})
gradient_checkpointing: bool = dataclasses.field(default=True, metadata={"help": "Use gradient checkpointing"})
dataset_text_field: str = dataclasses.field(default="text", metadata={"help": "Text field in the dataset"})
report_to: str = dataclasses.field(default="wandb", metadata={"help": "Where to report metrics"})
logging_steps: int = dataclasses.field(default=16, metadata={"help": "Log every N steps"})
evaluation_strategy: str = dataclasses.field(default="steps", metadata={"help": "Evaluation strategy"})
eval_steps: int = dataclasses.field(default=100, metadata={"help": "Evaluate every N steps"})
save_strategy: str = dataclasses.field(default="steps", metadata={"help": "Save strategy"})
save_steps: int = dataclasses.field(default=1500, metadata={"help": "Save every N steps"})
offload_optimizer: bool = dataclasses.field(default=False, metadata={"help": "Offload optimizer to CPU"})
@dataclasses.dataclass
class LoraArguments:
lora_r: int = 64
lora_alpha: int = 16
lora_dropout: float = 0.05
lora_target_modules: List[str] = dataclasses.field(default_factory=lambda: [])
lora_bias: str = "none"
q_lora: bool = False
def maybe_zero_3(param):
if hasattr(param, "ds_id"):
assert param.ds_status == ZeroParamStatus.NOT_AVAILABLE
with zero.GatheredParameters([param]):
param = param.data.detach().cpu().clone()
else:
param = param.detach().cpu().clone()
return param
def get_peft_state_maybe_zero_3(named_params, bias):
if bias == "none":
to_return = {k: t for k, t in named_params if "lora_" in k}
elif bias == "all":
to_return = {k: t for k, t in named_params if "lora_" in k or "bias" in k}
elif bias == "lora_only":
to_return = {}
maybe_lora_bias = {}
lora_bias_names = set()
for k, t in named_params:
if "lora_" in k:
to_return[k] = t
bias_name = k.split("lora_")[0] + "bias"
lora_bias_names.add(bias_name)
elif "bias" in k:
maybe_lora_bias[k] = t
for k, t in maybe_lora_bias.items():
if k in lora_bias_names:
to_return[k] = t
else:
raise NotImplementedError
to_return = {k: maybe_zero_3(v) for k, v in to_return.items()}
return to_return
def safe_save_model_for_hf_trainer(trainer: transformers.Trainer, output_dir: str, bias="none"):
is_zero3 = False
if hasattr(trainer.accelerator.state, "deepspeed_plugin"):
ds_plugin = trainer.accelerator.state.deepspeed_plugin
is_zero3 = ds_plugin.zero_stage == 3 if ds_plugin is not None else False
elif hasattr(trainer.args, "deepspeed") and trainer.args.deepspeed:
if os.path.exists(trainer.args.deepspeed):
with open(trainer.args.deepspeed, "r") as f:
ds_config = json.load(f)
is_zero3 = ds_config.get("zero_optimization", {}).get("stage", 0) == 3
if is_zero3:
state_dict = trainer.model_wrapped._zero3_consolidated_16bit_state_dict()
else:
if trainer.args.use_lora:
state_dict = get_peft_state_maybe_zero_3(trainer.model.named_parameters(), bias)
else:
state_dict = trainer.model.state_dict()
if trainer.args.should_save and trainer.args.local_rank == 0:
trainer._save(output_dir, state_dict=state_dict)
class DataCollatorForCompletionOnlyLM(DataCollatorForLanguageModeling):
def init(self, tokenizer, *args, **kwargs):
super().init(tokenizer, *args, **kwargs)
self.tokenizer = tokenizer
self.RESPONSE_KEY_NL = "### Response:\n"
self.response_token_ids = self.tokenizer.encode(self.RESPONSE_KEY_NL, add_special_tokens=False)
def torch_call(self, examples: List[Union[List[int], Any, Dict[str, Any]]]) -> Dict[str, Any]:
batch = super().torch_call(examples)
labels = batch["labels"].clone()
for i in range(len(examples)):
input_ids = batch["input_ids"][i]
response_start_idx = None
for idx in range(len(input_ids) - len(self.response_token_ids) + 1):
if torch.equal(input_ids[idx:idx + len(self.response_token_ids)], torch.tensor(self.response_token_ids)):
response_start_idx = idx + len(self.response_token_ids)
break
if response_start_idx is None:
labels[i, :] = -100
logger.warning(f"Could not find response key in example {i}, masking all labels")
else:
labels[i, :response_start_idx] = -100
batch["labels"] = labels
return batch
def create_peft_config(model_name_or_path, lora_args):
if "pythia" in model_name_or_path.lower():
target_modules = ["query_key_value"]
elif "qwen" in model_name_or_path.lower():
target_modules = ["q_proj", "k_proj", "v_proj"]
elif "llama" in model_name_or_path.lower():
target_modules = ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"]
elif "mistral" in model_name_or_path.lower():
target_modules = ["q_proj", "k_proj", "v_proj", "o_proj"]
else:
target_modules = ["q_proj", "k_proj", "v_proj", "o_proj"]
lora_args.lora_target_modules = target_modules
return LoraConfig(
task_type=TaskType.CAUSAL_LM,
r=lora_args.lora_r,
lora_alpha=lora_args.lora_alpha,
lora_dropout=lora_args.lora_dropout,
target_modules=target_modules,
bias=lora_args.lora_bias,
)
class CustomSFTTrainer(SFTTrainer):
def init(self, *args, **kwargs):
super().init(*args, **kwargs)
self.loss_log = []
self.step_counter = 0
def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None):
labels = inputs.get("labels")
outputs = model(**inputs)
logits = outputs.logits
if labels is None:
if return_outputs:
return outputs.loss, outputs
return outputs.loss
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()
loss_fct = torch.nn.CrossEntropyLoss(ignore_index=-100)
loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1))
if hasattr(self.args, 'local_rank') and self.args.local_rank == 0:
self.loss_log.append(loss.item())
log_dict = {
"mean_batch_loss": loss.item(),
"step": self.step_counter
}
try:
wandb.log(log_dict)
except Exception as e:
logger.warning(f"Error logging to wandb: {e}")
self.step_counter += 1
return (loss, outputs) if return_outputs else loss
def evaluate(self, *args, **kwargs):
try:
eval_result = super().evaluate(*args, **kwargs)
logger.info(f"Evaluation result: {eval_result}")
eval_loss = eval_result.get("eval_loss", None)
if eval_loss is not None:
perplexity = torch.exp(torch.tensor(eval_loss))
eval_result["perplexity"] = perplexity.item()
if self.args.local_rank == 0:
try:
wandb.log({
"eval_loss": eval_loss,
"perplexity": perplexity.item(),
"global_step": self.state.global_step,
"eval_step": self.state.global_step,
})
logger.info(f"Evaluation at step {self.state.global_step}: loss={eval_loss}, perplexity={perplexity.item()}")
except Exception as e:
logger.warning(f"Failed to log evaluation metrics to wandb: {e}")
else:
logger.warning("eval_loss not found in evaluation result")
return eval_result
except Exception as e:
logger.error(f"Error during evaluation: {e}")
return {"eval_loss": float("inf")}
def preprocess_batch(batch: Dict[str, List], tokenizer: AutoTokenizer, max_length: int) -> dict:
if "text" not in batch:
logger.error(f"Batch doesn't contain 'text' field: {batch.keys()}")
return {"input_ids": [], "attention_mask": []}
try:
return tokenizer(
batch["text"],
max_length=max_length,
truncation=True,
)
except Exception as e:
logger.error(f"Error tokenizing batch: {e}")
return {"input_ids": [], "attention_mask": []}
def train():
global local_rank
parser = HfArgumentParser((ModelArguments, DataArguments, CustomTrainingArguments, LoraArguments))
model_args, data_args, training_args, lora_args = parser.parse_args_into_dataclasses()
local_rank = training_args.local_rank
set_seed(42)
if torch.cuda.is_available():
torch.cuda.set_device(local_rank)
device_map = None
world_size = int(os.environ.get("WORLD_SIZE", 1))
ddp = world_size != 1
if training_args.deepspeed:
if local_rank == 0:
logger.info(f"Using DeepSpeed config: {training_args.deepspeed}")
if lora_args.q_lora:
device_map = {"": local_rank} if ddp else "auto"
if len(training_args.fsdp) > 0 or deepspeed.is_deepspeed_zero3_enabled():
logging.warning("FSDP or ZeRO3 are incompatible with QLoRA.")
if local_rank == 0:
try:
wandb.init(project="qwence_0422", name="qkv16_32")
logger.info("Wandb initialized successfully")
except Exception as e:
logger.error(f"Failed to initialize wandb: {e}")
wandb = None # Prevent further wandb logging attempts
logger.info(f"Loading model {model_args.model_name_or_path}")
model_load_kwargs = {}
if training_args.use_lora and lora_args.q_lora:
from transformers import BitsAndBytesConfig
model_load_kwargs["quantization_config"] = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True,
bnb_4bit_quant_type="nf4"
)
model_load_kwargs["device_map"] = device_map
try:
if training_args.deepspeed:
model = AutoModelForCausalLM.from_pretrained(
model_args.model_name_or_path,
trust_remote_code=True,
use_cache=not training_args.gradient_checkpointing,
torch_dtype=torch.bfloat16,
**model_load_kwargs
)
else:
model = AutoModelForCausalLM.from_pretrained(
model_args.model_name_or_path,
device_map=device_map,
trust_remote_code=True,
use_cache=not training_args.gradient_checkpointing,
torch_dtype=torch.bfloat16,
**model_load_kwargs
)
except Exception as e:
logger.error(f"Error loading model: {e}")
raise
try:
tokenizer = AutoTokenizer.from_pretrained(
model_args.model_name_or_path,
padding_side="right",
use_fast=True,
trust_remote_code=True,
)
except Exception as e:
logger.error(f"Error loading tokenizer: {e}")
raise
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
special_tokens = [END_KEY, INSTRUCTION_KEY, RESPONSE_KEY, RESPONSE_KEY_NL]
tokenizer.add_special_tokens({"additional_special_tokens": special_tokens})
if not training_args.use_lora:
model.resize_token_embeddings(len(tokenizer))
peft_config = None
if training_args.use_lora:
logger.info("Applying LoRA configuration")
try:
if lora_args.q_lora:
logger.info("Preparing model for k-bit training")
model = prepare_model_for_kbit_training(
model,
use_gradient_checkpointing=training_args.gradient_checkpointing
)
peft_config = create_peft_config(model_args.model_name_or_path, lora_args)
logger.info(f"LoRA target modules: {peft_config.target_modules}")
model = get_peft_model(model, peft_config)
torch.cuda.empty_cache()
if training_args.gradient_checkpointing:
model.enable_input_require_grads()
trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
total_params = sum(p.numel() for p in model.parameters())
param_efficiency = 100 * trainable_params / total_params if total_params > 0 else 0
logger.info(f"Total trainable params: {trainable_params:,} ({param_efficiency:.2f}% of total)")
logger.info(f"Total params: {total_params:,}")
except Exception as e:
logger.error(f"Error setting up LoRA: {e}")
raise
from datasets import load_dataset
logger.info(f"Loading dataset from {data_args.data_path}")
try:
dataset = load_dataset(data_args.data_path)["train"]
except Exception as e:
logger.error(f"Error loading dataset: {e}")
raise
if data_args.max_samples is not None and data_args.max_samples < len(dataset):
dataset = dataset.select(range(data_args.max_samples))
logger.info(f"Dataset loaded with {dataset.num_rows} examples")
def format_example(example):
context = example.get("context", example.get("input", None))
try:
if context:
text = PROMPT_WITH_INPUT_FORMAT.format(
instruction=example["instruction"],
input=context,
response=example["response"]
)
else:
text = PROMPT_NO_INPUT_FORMAT.format(
instruction=example["instruction"],
response=example["response"]
)
return {"text": text}
except Exception as e:
logger.error(f"Error formatting example: {e}")
return {"text": "Error processing this example."}
start_time = time.time()
dataset = dataset.map(format_example)
preprocessing_time = time.time() - start_time
logger.info(f"Dataset formatting took {preprocessing_time:.2f} seconds")
logger.info("Preprocessing dataset")
_preprocessing_function = partial(preprocess_batch, tokenizer=tokenizer, max_length=training_args.max_seq_length)
processed_dataset = dataset.shuffle(seed=42)
try:
split_dataset = processed_dataset.train_test_split(test_size=data_args.test_size, seed=42)
start_time = time.time()
remove_columns = []
for col in ["instruction", "context", "response", "text", "category"]:
if col in split_dataset["train"].column_names:
remove_columns.append(col)
train_dataset = split_dataset["train"].map(
_preprocessing_function,
batched=True,
remove_columns=remove_columns,
num_proc=4,
)
eval_dataset = split_dataset["test"].map(
_preprocessing_function,
batched=True,
remove_columns=remove_columns,
num_proc=4,
)
except Exception as e:
logger.error(f"Error preprocessing dataset: {e}")
raise
logger.info(f"Train dataset has {train_dataset.num_rows} rows before filtering")
max_length = training_args.max_seq_length
try:
def is_valid_length(example):
return len(example["input_ids"]) < max_length
train_dataset = train_dataset.filter(is_valid_length)
logger.info(f"Train dataset has {train_dataset.num_rows} rows after filtering for truncated records")
logger.info(f"Eval dataset has {eval_dataset.num_rows} rows before filtering")
eval_dataset = eval_dataset.filter(is_valid_length)
logger.info(f"Eval dataset has {eval_dataset.num_rows} rows after filtering for truncated records")
if train_dataset.num_rows == 0:
raise ValueError("Filtered training dataset is empty! Check max_seq_length setting.")
if eval_dataset.num_rows == 0:
logger.warning("Filtered evaluation dataset is empty, using a subset of the training data instead")
eval_dataset = train_dataset.select(range(min(100, train_dataset.num_rows)))
except Exception as e:
logger.error(f"Error filtering dataset: {e}")
raise
logger.info(f"Evaluation dataset size: {eval_dataset.num_rows}") # Log eval dataset size
data_collator = DataCollatorForCompletionOnlyLM(
tokenizer=tokenizer,
mlm=False,
return_tensors="pt",
pad_to_multiple_of=8
)
if train_dataset.num_rows == 0:
raise ValueError("Training dataset is empty after preprocessing!")
if dist.is_initialized():
local_size = len(train_dataset) // dist.get_world_size()
all_sizes = [torch.tensor(0, device=f"cuda:{local_rank}") for _ in range(dist.get_world_size())]
local_tensor = torch.tensor(local_size, device=f"cuda:{local_rank}")
dist.all_gather(all_sizes, local_tensor)
if local_rank == 0:
logger.info(f"Data distribution across processes: {[t.item() for t in all_sizes]}")
total_examples = len(train_dataset)
logger.info(f"Total examples (corrected): {total_examples}")
trainer = CustomSFTTrainer(
model=model,
tokenizer=tokenizer,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
data_collator=data_collator,
peft_config=peft_config if training_args.use_lora else None,
packing=training_args.packing,
)
if local_rank == 0:
logger.info(f"Optimizer: {trainer.args.optim}")
logger.info(f"Learning rate: {trainer.args.learning_rate}")
logger.info(f"Batch size per device: {trainer.args.per_device_train_batch_size}")
logger.info(f"Gradient accumulation steps: {trainer.args.gradient_accumulation_steps}")
logger.info(f"Total batch size: {trainer.args.per_device_train_batch_size * trainer.args.gradient_accumulation_steps * world_size}")
start_train_time = time.time()
if local_rank == 0:
logger.info("Starting training")
try:
original_training_step = trainer.training_step
def training_step_with_timeout(*args, **kwargs):
step_start = time.time()
result = original_training_step(*args, **kwargs)
step_duration = time.time() - step_start
if step_duration > 30:
logger.warning(f"Step took {step_duration:.2f}s, which is unusually long")
return result
trainer.training_step = training_step_with_timeout
train_result = trainer.train()
except Exception as e:
logger.error(f"Error during training: {e}")
if local_rank == 0:
logger.warning("Training failed, attempting to save current model state")
try:
safe_save_model_for_hf_trainer(trainer=trainer, output_dir=f"{training_args.output_dir}_error_recovery", bias=lora_args.lora_bias)
except Exception as es:
logger.error(f"Could not save model after error: {es}")
raise
train_duration = time.time() - start_train_time
if local_rank == 0:
logger.info(f"Training completed in {train_duration:.2f} seconds")
if wandb.run is not None:
try:
wandb.log({"total_train_time": train_duration})
except:
logger.warning("Failed to log training duration to wandb")
if local_rank == 0:
logger.info(f"Saving model to {training_args.output_dir}")
try:
safe_save_model_for_hf_trainer(trainer=trainer, output_dir=training_args.output_dir, bias=lora_args.lora_bias)
tokenizer.save_pretrained(training_args.output_dir)
training_config = {
"model": {k: v for k, v in vars(model_args).items() if not k.startswith('_')},
"training": {
k: str(v) if not isinstance(v, (int, float, str, bool, list, dict, type(None)))
else v for k, v in vars(training_args).items()
if not (k.startswith('_') or callable(v))
},
"lora": {k: v for k, v in vars(lora_args).items() if not k.startswith('_')},
"data": {k: v for k, v in vars(data_args).items() if not k.startswith('_')},
"train_time": train_duration,
"num_examples": len(train_dataset),
"world_size": world_size
}
with open(os.path.join(training_args.output_dir, "training_config.json"), "w") as f:
json.dump(training_config, f, indent=2)
except Exception as e:
logger.error(f"Error saving model: {e}")
del model
del trainer
torch.cuda.empty_cache()
gc.collect()
return train_result
if name == "main":
try:
train()
except Exception as e:
logger.error(f"Training failed with error: {e}")
import traceback
logger.error(traceback.format_exc())
import sys
sys.exit(1)
'''
CUDA_VISIBLE_DEVICES=0,1 torchrun --nproc_per_node=2 --nnodes=1 --node_rank=0 --master_port=6002 celora.py \
Expected behavior
The gpu memory keeps the same usage
ds_report output
DeepSpeed C++/CUDA extension op report
[WARNING] async_io requires the dev libaio .so object and headers but these were not found.
[WARNING] async_io: please install the libaio-dev package with apt
[WARNING] If libaio is already installed (perhaps from source), try setting the CFLAGS and LDFLAGS environment variables to where it can be found.
async_io ............... [NO] ....... [NO]
fused_adam ............. [NO] ....... [OKAY]
cpu_adam ............... [NO] ....... [OKAY]
cpu_adagrad ............ [NO] ....... [OKAY]
cpu_lion ............... [NO] ....... [OKAY]
[WARNING] Please specify the CUTLASS repo directory as environment variable $CUTLASS_PATH
evoformer_attn ......... [NO] ....... [NO]
[WARNING] FP Quantizer is using an untested triton version (3.2.0), only 2.3.(0, 1) and 3.0.0 are known to be compatible with these kernels
fp_quantizer ........... [NO] ....... [NO]
fused_lamb ............. [NO] ....... [OKAY]
fused_lion ............. [NO] ....... [OKAY]
gds .................... [NO] ....... [NO]
transformer_inference .. [NO] ....... [OKAY]
inference_core_ops ..... [NO] ....... [OKAY]
cutlass_ops ............ [NO] ....... [OKAY]
quantizer .............. [NO] ....... [OKAY]
ragged_device_ops ...... [NO] ....... [OKAY]
ragged_ops ............. [NO] ....... [OKAY]
random_ltd ............. [NO] ....... [OKAY]
[WARNING] sparse_attn requires a torch version >= 1.5 and < 2.0 but detected 2.6
[WARNING] using untested triton version (3.2.0), only 1.0.0 is known to be compatible
sparse_attn ............ [NO] ....... [NO]
spatial_inference ...... [NO] ....... [OKAY]
transformer ............ [NO] ....... [OKAY]
stochastic_transformer . [NO] ....... [OKAY]
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start with the inline script's CustomSFTTrainer.evaluate path and its torchrun/DeepSpeed stage 2 configuration, then reproduce the reported GPU 0/GPU 1 memory difference during evaluation and training. Compare per-rank data sampling and memory usage to determine whether the imbalance is reproducible; done means identifying a confirmed DeepSpeed or SFTTrainer cause and a regression test or documented reproduction.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- distributed-systems, machine-learning
- Issue type
- Bug
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Stale
- Clarity
- Needs clarification
- Newbie friendliness
- 25/100