PatchDPOTrainer

This commit is contained in:
Daniel Han-Chen 2023-12-31 18:36:43 +11:00
commit 148e7db6c3
3 changed files with 8 additions and 17 deletions

View file

@ -12,6 +12,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from .loader import FastLanguageModel, FastDPOTrainer
from .loader import FastLanguageModel
from .llama import FastLlamaModel
from .mistral import FastMistralModel
from .dpo import PatchDPOTrainer

View file

@ -17,9 +17,6 @@ from transformers.utils.notebook import (
NotebookTrainingTracker,
NotebookProgressCallback,
)
from transformers.trainer import DEFAULT_PROGRESS_CALLBACK
from trl import DPOTrainer
import types
DPOTrainer_metrics = [
"rewards/chosen",
@ -32,6 +29,7 @@ DPOTrainer_metrics = [
"logits/chosen",
]
def NotebookProgressCallback_on_train_begin(self, args, state, control, **kwargs):
self.first_column = "Epoch" if args.evaluation_strategy == IntervalStrategy.EPOCH else "Step"
self.training_loss = 0
@ -67,17 +65,10 @@ def NotebookProgressCallback_on_log(self, args, state, control, logs=None, **kwa
pass
class FastDPOTrainer(DPOTrainer):
def PatchDPOTrainer():
# Patch DPO notebook printing
if (DEFAULT_PROGRESS_CALLBACK is NotebookProgressCallback):
DEFAULT_PROGRESS_CALLBACK.on_train_begin = types.MethodType(
NotebookProgressCallback_on_train_begin,
DEFAULT_PROGRESS_CALLBACK,
)
DEFAULT_PROGRESS_CALLBACK.on_log = types.MethodType(
NotebookProgressCallback_on_log,
DEFAULT_PROGRESS_CALLBACK,
)
pass
from transformers.trainer import DEFAULT_PROGRESS_CALLBACK
DEFAULT_PROGRESS_CALLBACK.on_train_begin = NotebookProgressCallback_on_train_begin
DEFAULT_PROGRESS_CALLBACK.on_log = NotebookProgressCallback_on_log
pass

View file

@ -14,7 +14,6 @@
from .llama import FastLlamaModel, logger
from .mistral import FastMistralModel
from .dpo import FastDPOTrainer
from transformers import AutoConfig
from transformers import __version__ as transformers_version