fixing the chatml None error
This commit is contained in:
parent
47fc79df6d
commit
6e8e70c987
1 changed files with 59 additions and 68 deletions
|
|
@ -357,48 +357,25 @@ def format_dataset(
|
|||
conversations = []
|
||||
num_examples = len(examples[list(examples.keys())[0]])
|
||||
|
||||
# NEW: Check if this is user-provided or auto-detected
|
||||
is_user_provided = custom_format_mapping is not None # Passed explicitly
|
||||
|
||||
# Preserve non-mapped columns ONLY if auto-detected
|
||||
preserved_columns = {}
|
||||
if not is_user_provided: # Only preserve for auto-detection
|
||||
all_columns = set(examples.keys())
|
||||
mapped_columns = set(custom_mapping.keys())
|
||||
non_mapped_columns = all_columns - mapped_columns
|
||||
|
||||
for col in non_mapped_columns:
|
||||
preserved_columns[col] = examples[col]
|
||||
# Preserve non-mapped columns
|
||||
all_columns = set(examples.keys())
|
||||
mapped_columns = set(custom_mapping.keys())
|
||||
preserved_columns = {
|
||||
col: examples[col]
|
||||
for col in all_columns - mapped_columns
|
||||
}
|
||||
|
||||
for i in range(num_examples):
|
||||
convo = []
|
||||
|
||||
# Enforce standard role order
|
||||
role_order = ['system', 'user', 'assistant']
|
||||
|
||||
for target_role in role_order:
|
||||
for target_role in ['system', 'user', 'assistant']:
|
||||
for col_name, role in custom_mapping.items():
|
||||
if role == target_role and col_name in examples:
|
||||
content = examples[col_name][i]
|
||||
|
||||
# NEW: Different behavior based on mapping source
|
||||
if is_user_provided:
|
||||
# User explicitly mapped this - always include even if empty
|
||||
convo.append({"role": role, "content": str(content) if content else ""})
|
||||
else:
|
||||
# Auto-detected - skip empty (original behavior)
|
||||
if content and str(content).strip():
|
||||
convo.append({"role": role, "content": str(content)})
|
||||
|
||||
if content and str(content).strip():
|
||||
convo.append({"role": role, "content": str(content)})
|
||||
conversations.append(convo)
|
||||
|
||||
result = {"conversations": conversations}
|
||||
|
||||
# Only add preserved columns if auto-detected
|
||||
if not is_user_provided:
|
||||
result.update(preserved_columns)
|
||||
|
||||
return result
|
||||
return {"conversations": conversations, **preserved_columns}
|
||||
|
||||
|
||||
try:
|
||||
|
|
@ -556,36 +533,38 @@ def format_dataset(
|
|||
|
||||
else:
|
||||
warnings.append(f"Unknown format, attempting standardization")
|
||||
try:
|
||||
standardized = standardize_chat_format(
|
||||
dataset, tokenizer, aliases_for_system,
|
||||
aliases_for_user, aliases_for_assistant,
|
||||
batch_size, num_proc
|
||||
)
|
||||
return {
|
||||
"dataset": standardized,
|
||||
"detected_format": "unknown",
|
||||
"final_format": f"chatml_{detected['chat_column']}",
|
||||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": warnings
|
||||
}
|
||||
except Exception as e:
|
||||
warnings.append(f"Standardization failed: {e}")
|
||||
return {
|
||||
"dataset": dataset,
|
||||
"detected_format": "unknown",
|
||||
"final_format": "unknown",
|
||||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": False,
|
||||
"requires_manual_mapping": True,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": warnings
|
||||
}
|
||||
if detected["chat_column"]:
|
||||
try:
|
||||
standardized = standardize_chat_format(
|
||||
dataset, tokenizer, aliases_for_system,
|
||||
aliases_for_user, aliases_for_assistant,
|
||||
batch_size, num_proc
|
||||
)
|
||||
return {
|
||||
"dataset": standardized,
|
||||
"detected_format": "unknown",
|
||||
"final_format": f"chatml_{detected['chat_column']}",
|
||||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": True,
|
||||
"requires_manual_mapping": False,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": warnings
|
||||
}
|
||||
except Exception as e:
|
||||
warnings.append(f"Standardization failed: {e}")
|
||||
|
||||
return {
|
||||
"dataset": dataset,
|
||||
"detected_format": "unknown",
|
||||
"final_format": "unknown",
|
||||
"chat_column": detected["chat_column"],
|
||||
"is_standardized": False,
|
||||
"requires_manual_mapping": True,
|
||||
"is_multimodal": multimodal_info["is_multimodal"],
|
||||
"multimodal_info": multimodal_info,
|
||||
"warnings": warnings
|
||||
}
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unknown format_type: {format_type}")
|
||||
|
|
@ -816,8 +795,10 @@ def format_and_template_dataset(
|
|||
)
|
||||
|
||||
# Step 2: Apply chat template
|
||||
if "gemma" in model_name.lower() and not dataset_info["is_multimodal"] and (format_type != "alpaca" or (format_type == "auto" and dataset_info["detected_format"] != "alpaca")):
|
||||
print("remove_bos_prefix is true")
|
||||
# Gemma emits a leading <bos> that must be stripped for text-only chatml/sharegpt.
|
||||
is_alpaca = format_type == "alpaca" or (format_type == "auto" and dataset_info["detected_format"] == "alpaca")
|
||||
is_gemma = "gemma" in model_name.lower()
|
||||
if is_gemma and not dataset_info["is_multimodal"] and not is_alpaca:
|
||||
remove_bos_prefix = True
|
||||
template_result = apply_chat_template_to_dataset(
|
||||
dataset_info=dataset_info,
|
||||
|
|
@ -839,14 +820,24 @@ def format_and_template_dataset(
|
|||
all_warnings = dataset_info.get("warnings", []) + template_result.get("warnings", [])
|
||||
all_errors = template_result.get("errors", [])
|
||||
|
||||
# If format_dataset returned "unknown" but apply_chat_template rescued
|
||||
# it via heuristic detection, update final_format to reflect reality.
|
||||
final_format = dataset_info["final_format"]
|
||||
requires_manual = dataset_info.get("requires_manual_mapping", False)
|
||||
if final_format == "unknown" and template_result["success"]:
|
||||
out_ds = template_result["dataset"]
|
||||
if hasattr(out_ds, "column_names") and "text" in out_ds.column_names:
|
||||
final_format = "chatml_conversations"
|
||||
requires_manual = False
|
||||
|
||||
return {
|
||||
"dataset": template_result["dataset"],
|
||||
"detected_format": dataset_info["detected_format"],
|
||||
"final_format": dataset_info["final_format"],
|
||||
"final_format": final_format,
|
||||
"chat_column": dataset_info.get("chat_column"),
|
||||
"is_vlm": False, # This is LLM flow
|
||||
"success": template_result["success"],
|
||||
"requires_manual_mapping": dataset_info.get("requires_manual_mapping", False),
|
||||
"requires_manual_mapping": requires_manual,
|
||||
"warnings": all_warnings,
|
||||
"errors": all_errors,
|
||||
"summary": summary,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue