deepspeedai / deepspeedai/DeepSpeed

[BUG]

Open
#7,235 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

bug training
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

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.