Business
Jobs
  • About Us
  • Solutions
    • Job Postings
      Post your job and receive qualified candidates in 48h.
    • Candidate Assessments
      500+ technical and psychological tests, plus anti-fraud.
    • Headhunting
      Tailor-made executive search from start to finish.
    • Payroll + EOR
      Payroll dispersal and EOR across 15+ LATAM countries.
  • Pricing
  • Jobs

0

302
Views
Detención temprana en instancias de Bert Trainer

Estoy ajustando un modelo BERT para una tarea de clasificación multiclase. Mi problema es que no sé cómo agregar "detención anticipada" a esas instancias de Entrenador. ¿Algunas ideas?

over 4 years ago · Santiago Trujillo
1 answers
Answer question

0

Hay un par de cosas que debe hacer antes de usar correctamente EarlyStoppingCallback()

 from transformers import EarlyStoppingCallback ... ... # Defining the TrainingArguments() arguments args = TrainingArguments( f"training_with_callbacks", evaluation_strategy ='steps', eval_steps = 50, # Evaluation and Save happens every 50 steps save_total_limit = 5, # Only last 5 models are saved. Older ones are deleted. learning_rate=2e-5, per_device_train_batch_size=batch_size, per_device_eval_batch_size=batch_size, num_train_epochs=5, weight_decay=0.01, push_to_hub=False, metric_for_best_model = 'f1', load_best_model_at_end=True)

Necesitas:

  1. Use load_best_model_at_end = True ( EarlyStoppingCallback() requiere que esto sea True ).
  2. evaluation_strategy = 'steps' en lugar de 'epoch' .
  3. eval_steps = 50 (evalúa las métricas después de N pasos).
  4. metric_for_best_model = 'f1' ,

En tu Trainer() :

 trainer = Trainer( model, args, ... compute_metrics=compute_metrics, callbacks = [EarlyStoppingCallback(early_stopping_patience=3)] )

Por supuesto, cuando usas compute_metrics , por ejemplo, puede ser una función como:

 def compute_metrics(p): pred, labels = p pred = np.argmax(pred, axis=1) accuracy = accuracy_score(y_true=labels, y_pred=pred) recall = recall_score(y_true=labels, y_pred=pred) precision = precision_score(y_true=labels, y_pred=pred) f1 = f1_score(y_true=labels, y_pred=pred) return {"accuracy": accuracy, "precision": precision, "recall": recall, "f1": f1}

El retorno de compute_metrics() debe ser un diccionario y puede acceder a cualquier métrica que desee/calcular dentro de la función y regresar.

over 4 years ago · Santiago Trujillo Report
Answer question
Find remote jobs

Discover the new way to find a job!

Top jobs
Top job categories
Business
Post vacancy Pricing Sales
Legal
Terms and conditions Privacy policy
© 2026 PeakU Inc. All Rights Reserved.
Andres GPT
Show me some job opportunities
There's an error!