Precision issues

This commit is contained in:
Daniel Han 2025-03-14 08:33:33 -07:00
commit 8cfe8a57e6
2 changed files with 3 additions and 4 deletions

View file

@ -12,7 +12,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
__version__ = "2025.3.12"
__version__ = "2025.3.13"
__all__ = [
"SUPPORTS_BFLOAT16",

View file

@ -238,9 +238,8 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
"use_fp16 = getattr(args, 'fp16', False)\n"\
"force_float32 = False\n"\
"if os.environ.get('UNSLOTH_FORCE_FLOAT32', '0') == '1':\n"\
" if use_bf16 or use_fp16:\n"\
" print('Unsloth: Switching to float32 training since model cannot work with float16')\n"\
" force_float32 = True\n"\
" print('Unsloth: Switching to float32 training since model cannot work with float16')\n"\
" force_float32 = True\n"\
"mixed_precision_dtype = os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32')\n"\
"dtype = getattr(model.config, 'torch_dtype', None)\n"\
"if dtype is None: dtype = model.get_input_embeddings().dtype\n"\