Skip to content
Open
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
2 changes: 1 addition & 1 deletion angelspec/data/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -181,7 +181,7 @@ def load_conversation_dataset(args):
drop_overlength_flag = getattr(args, "drop_overlength", False)
cache_params = (
f"{dataset_name}-{args.train_data_path}{file_stat}-{args.target_model_path}"
f"-{max_length}-{chat_template_name}-ltlo={last_turn_loss_only_flag}"
f"-{max_length}-{chat_template_name}-{prompt_key}-ltlo={last_turn_loss_only_flag}"
f"-defer={defer_tokenization}-decode={train_with_decode}"
f"-mlt={min_loss_tokens_val}-drop={drop_overlength_flag}"
)
Expand Down
4 changes: 2 additions & 2 deletions angelspec/data/parse.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,7 +101,7 @@ def _prepare_text(self, conversation: "Conversation", preformatted: bool, **kwar
return self.format(conversation, **kwargs)

def _ensure_pad_token(self):
if not self.tokenizer.pad_token_id:
if self.tokenizer.pad_token_id is None:
self.tokenizer.pad_token_id = self.tokenizer.unk_token_id

def _tokenize_with_loss_mask(
Expand Down Expand Up @@ -254,7 +254,7 @@ def parse(
assistant_pattern = (
re.escape(self.assistant_message_separator)
+ r"([\s\S]*?(?:"
+ re.escape(self.chat_template.end_of_turn_token)
+ re.escape(self.chat_template.end_of_turn_token or "")
+ "|$))"
)
return self._tokenize_with_loss_mask(
Expand Down
16 changes: 14 additions & 2 deletions angelspec/data/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -832,6 +832,16 @@ def load_hf_dataset(data_path: str):
load_local_json, gen_kwargs={"data_path": data_path}
)
ext = os.path.splitext(data_path)[1].lower()
if ext in (".csv", ".tsv"):
return load_dataset(
"csv",
data_files=data_path,
sep="\t" if ext == ".tsv" else ",",
split="train",
streaming=True,
)
if ext == ".txt":
return load_dataset("text", data_files=data_path, split="train", streaming=True)
fmt = {".parquet": "parquet", ".arrow": "arrow"}.get(ext, "json")
return load_dataset(fmt, data_files=data_path, split="train", streaming=True)

Expand Down Expand Up @@ -863,8 +873,10 @@ def load_hf_dataset(data_path: str):
raise FileNotFoundError(f"Local dataset path not found: {data_path}")

# hub path — try native load_dataset first (handles Arrow, Parquet, etc.),
# fall back to manual JSON download for repos with mixed-type columns
_KEEP_COLUMNS = frozenset({"id", "conversations", "text", "messages"})
# fall back to manual JSON download for repos with mixed-type columns.
# Keep the per-sample top-level fields the chat template consumes
# (dataset.py threads them into parser.format).
_KEEP_COLUMNS = frozenset({"id", "conversations", "text", "messages", "tools", "reasoning_effort"})
try:
ds = load_dataset(data_path, split="train", streaming=True)
drop_cols = [c for c in (ds.column_names or []) if c not in _KEEP_COLUMNS]
Expand Down
3 changes: 3 additions & 0 deletions angelspec/models/draft/deepseek_eagle.py
Original file line number Diff line number Diff line change
Expand Up @@ -181,12 +181,14 @@ def _init_rope(self):
self.rotary_emb = LlamaLinearScalingRotaryEmbedding(
rope_dim,
max_position_embeddings=self.max_position_embeddings,
base=rope_theta,
scaling_factor=scaling_factor,
)
elif scaling_type == "dynamic":
self.rotary_emb = LlamaDynamicNTKScalingRotaryEmbedding(
rope_dim,
max_position_embeddings=self.max_position_embeddings,
base=rope_theta,
scaling_factor=scaling_factor,
)
elif scaling_type == "llama3":
Expand All @@ -203,6 +205,7 @@ def _init_rope(self):
self.rotary_emb = LlamaYarnRotaryEmbedding(
rope_dim,
max_position_embeddings=self.max_position_embeddings,
base=rope_theta,
original_max_position_embeddings=rget("original_max_position_embeddings"),
scaling_factor=scaling_factor,
beta_fast=rget("beta_fast"),
Expand Down
2 changes: 2 additions & 0 deletions angelspec/models/draft/llama3_eagle.py
Original file line number Diff line number Diff line change
Expand Up @@ -1176,6 +1176,7 @@ def rope_get(key, default=None):
self.rotary_emb = LlamaLinearScalingRotaryEmbedding(
self.head_dim,
max_position_embeddings=self.max_position_embeddings,
base=getattr(self.config, "rope_theta", 10000),
scaling_factor=scaling_factor,
)
elif scaling_type == "dynamic":
Expand All @@ -1186,6 +1187,7 @@ def rope_get(key, default=None):
self.rotary_emb = LlamaDynamicNTKScalingRotaryEmbedding(
self.head_dim,
max_position_embeddings=self.max_position_embeddings,
base=getattr(self.config, "rope_theta", 10000),
scaling_factor=scaling_factor,
)
elif scaling_type == "llama3":
Expand Down
2 changes: 1 addition & 1 deletion angelspec/transfer/mooncake/eagle_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -377,7 +377,7 @@ def get(
(
"last_hidden_states",
shapes["last_hidden_states"],
dtypes.get("hidden_states", HIDDEN_STATES_STORAGE_DTYPE),
dtypes.get("last_hidden_states", HIDDEN_STATES_STORAGE_DTYPE),
)
)

Expand Down
3 changes: 2 additions & 1 deletion angelspec/utils/logging.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,8 +34,9 @@ def _get_logger_level():
level_str = os.getenv("ANGELSPEC_LOG_LEVEL", "INFO").upper()
try:
log_level = getattr(logging, level_str)
except ValueError:
except AttributeError:
logging.warning("Invalid log level: %s, defaulting to WARNING", level_str)
log_level = logging.WARNING
return log_level


Expand Down
6 changes: 6 additions & 0 deletions angelspec/utils/profiling.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,12 @@ def _profile_simple_loop(iterator, args, name):


def _create_torch_profiler(args, name):
if args.profile_step_end <= args.profile_step_start:
raise ValueError(
f"profile_step_end ({args.profile_step_end}) must be greater than "
f"profile_step_start ({args.profile_step_start}) when use_pytorch_profiler "
"is enabled."
)
return torch.profiler.profile(
schedule=torch.profiler.schedule(
wait=max(args.profile_step_start - 1, 0),
Expand Down
4 changes: 2 additions & 2 deletions angelspec/utils/usp.py
Original file line number Diff line number Diff line change
Expand Up @@ -125,8 +125,8 @@ def _slice_and_pad(tensor: torch.Tensor, axis: int, pad_value: int = 0):
)
attention_mask[:, :valid_len] = 1

usp_chunk_size = max(local_len - ttt_length, 0)
ring_chunk = usp_chunk_size * sp_ulysses_size
chunk_len = max(local_len - ttt_length, 0)
ring_chunk = chunk_len * sp_ulysses_size
ring_start = ring_rank * ring_chunk
position_ids = torch.arange(
ring_start, ring_start + ring_chunk, device=input_ids.device, dtype=torch.long
Expand Down
2 changes: 1 addition & 1 deletion examples/hy3-dfly/run.sh
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@ echo "=============================================="
python3 -m angelspec.train_entry \
--config "$CONFIG_FILE" \
training.training_num_gpus_per_node=4 \
training.num_nodes="$NUM_NODES" \
training.training_num_nodes="$NUM_NODES" \
inference.inference_num_gpus="$INFERENCE_GPUS" \
inference.inference_num_gpus_per_engine=4 \
inference.inference_num_gpus_per_node="$GPUS_PER_NODE" \
Expand Down
2 changes: 1 addition & 1 deletion examples/hy3-mtp/run.sh
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@ echo "=============================================="
python3 -m angelspec.train_entry \
--config "$CONFIG_FILE" \
training.training_num_gpus_per_node=4 \
training.num_nodes="$NUM_NODES" \
training.training_num_nodes="$NUM_NODES" \
training.attention_backend=usp \
training.sp_ulysses_size=4 \
inference.inference_num_gpus="$INFERENCE_GPUS" \
Expand Down
6 changes: 6 additions & 0 deletions tools/generate_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -165,6 +165,7 @@ def call_sglang(
messages = data["conversations"]
regenerated_messages = []
total_output_tokens = 0
resp = None

if messages[0]["role"] == "assistant":
data["status"] = "error"
Expand Down Expand Up @@ -199,6 +200,11 @@ def call_sglang(
data["error"] = f"Invalid message role: {message['role']}"
return data

if resp is None:
data["status"] = "error"
data["error"] = "No user message in conversation"
return data

data["output_tokens"] = total_output_tokens
data["input_tokens"] = resp.usage.prompt_tokens
data["context_length"] = data["input_tokens"] + total_output_tokens
Expand Down