diff --git a/configs/default.yaml b/configs/default.yaml index a44d1e2..cbef25e 100644 --- a/configs/default.yaml +++ b/configs/default.yaml @@ -13,8 +13,12 @@ use_liger_kernel: true remove_unused_columns: false optim: schedule_free_adamw -learning_rate: 0.00001 +learning_rate: 0.0001 neftune_noise_alpha: 5 +label_smoothing_factor: 0.1 +weight_decay: 0.01 +warmup_ratio: 0.1 + # LoRA diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index 04ed5a3..dba475c 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -204,9 +204,10 @@ def main(): def validate_columns(tokenized_ds): - ref_cols = set( - ["input_ids", "attention_mask", "labels", "ctx_features", "ctx_attn_mask"] - ) + cols = ["input_ids", "attention_mask", "labels"] + if "ctx_features" in tokenized_ds["train"].column_names: + cols += ["ctx_features", "ctx_attn_mask"] + ref_cols = set(cols) assert ( set(tokenized_ds["train"].column_names) == ref_cols ), f"Columns mismatch: {set(tokenized_ds['train'].column_names)} != {ref_cols}"