调用keras_nlp.from_preset触发Segmentation fault问题求助
Keras-NLP使用PyTorch后端初始化DebertaV3Preprocessor时出现Segmentation fault (core dumped)
我在使用keras_nlp库时遇到了Segmentation fault (core dumped)错误,错误发生在初始化DebertaV3Preprocessor的步骤。我的运行环境是搭载80GB显存的Nvidia-A100显卡,显存充足,排除显存不足的可能。
相关代码
import os os.environ["KERAS_BACKEND"] = "torch" # "jax" or "tensorflow" or "torch" import keras_nlp import keras_core as keras import keras_core.backend as K import torch import tensorflow as tf import numpy as np import pandas as pd import matplotlib.pyplot as plt import matplotlib as mpl cmap = mpl.cm.get_cmap('coolwarm') class CFG: verbose = 0 # Verbosity wandb = True # Weights & Biases logging competition = 'llm-detect-ai-generated-text' # Competition name _wandb_kernel = 'awsaf49' # WandB kernel comment = 'DebertaV3-MaxSeq_200-ext_s-torch' # Comment description preset = "deberta_v3_base_en" # Name of pretrained models sequence_length = 200 # Input sequence length device = 'TPU' # Device seed = 42 # Random seed num_folds = 5 # Total folds selected_folds = [0, 1, 2] # Folds to train on epochs = 3 # Training epochs batch_size = 3 # Batch size drop_remainder = True # Drop incomplete batches cache = True # Caches data after one iteration, use only with `TPU` to avoid OOM scheduler = 'cosine' # Learning rate scheduler class_names = ["real", "fake"] # Class names [A, B, C, D, E] num_classes = len(class_names) # Number of classes class_labels = list(range(num_classes)) # Class labels [0, 1, 2, 3, 4] label2name = dict(zip(class_labels, class_names)) # Label to class name mapping name2label = {v: k for k, v in label2name.items()} # Class name to label mapping keras.utils.set_random_seed(CFG.seed) def get_device(): "Detect and intializes GPU/TPU automatically" try: # Connect to TPU tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect() # Set TPU strategy strategy = tf.distribute.TPUStrategy(tpu) print(f'> Running on TPU', tpu.master(), end=' | ') print('Num of TPUs: ', strategy.num_replicas_in_sync) device=CFG.device except: # If TPU is not available, detect GPUs gpus = tf.config.list_logical_devices('GPU') ngpu = len(gpus) # Check number of GPUs if ngpu: # Set GPU strategy strategy = tf.distribute.MirroredStrategy(gpus) # single-GPU or multi-GPU # Print GPU details print("> Running on GPU", end=' | ') print("Num of GPUs: ", ngpu) device='GPU' else: # If no GPUs are available, use CPU print("> Running on CPU") strategy = tf.distribute.get_strategy() device='CPU' return strategy, device # Initialize GPU/TPU/TPU-VM strategy, CFG.device = get_device() CFG.replicas = strategy.num_replicas_in_sync BASE_PATH = '/some/path/' print(1) preprocessor = keras_nlp.models.DebertaV3Preprocessor.from_preset( preset=CFG.preset, # Name of the model sequence_length=CFG.sequence_length, # Max sequence length, will be padded if shorter ) print(2)
完整运行日志
$python test.py Using PyTorch backend. /mypath/test.py:22: MatplotlibDeprecationWarning: The get_cmap function was deprecated in Matplotlib 3.7 and will be removed two minor releases later. Use ``matplotlib.colormaps[name]`` or ``matplotlib.colormaps.get_cmap(obj)`` instead. cmap = mpl.cm.get_cmap('coolwarm') > Running on GPU | Num of GPUs: 1 1 Segmentation fault (core dumped)
可能的问题原因及解决方案
1. TensorFlow与PyTorch后端冲突
代码中同时导入了TensorFlow和PyTorch,并且使用了TensorFlow的分布式策略(MirroredStrategy),但Keras后端设置为PyTorch,这会导致底层设备管理逻辑冲突,触发段错误。
解决办法:
- 移除所有TensorFlow相关的导入和设备检测代码,改用PyTorch原生的设备管理逻辑。修改后的
get_device函数示例:
def get_device(): if torch.cuda.is_available(): device = torch.device("cuda") print(f'> Running on GPU | Num of GPUs: {torch.cuda.device_count()}') else: device = torch.device("cpu") print('> Running on CPU') return device
- 删除
import tensorflow as tf语句,以及CFG中与TPU相关的配置(因为PyTorch后端不支持TF的TPU策略)。
2. 库版本不兼容
Keras-NLP、Keras-Core和PyTorch的版本不匹配可能导致底层代码崩溃。
解决办法:
- 升级到最新稳定版的相关库:
pip install --upgrade keras-core keras-nlp torch
3. 预训练模型加载异常
自动下载预训练模型的词汇表或配置文件时,可能出现文件损坏、权限问题或内存临时占用过高的情况。
解决办法:
- 手动下载DebertaV3的预训练配置和词汇表文件,保存到本地路径,然后通过
from_preset的cache_dir参数指定本地路径加载:
preprocessor = keras_nlp.models.DebertaV3Preprocessor.from_preset( preset=CFG.preset, sequence_length=CFG.sequence_length, cache_dir="/path/to/local/cache" )
内容的提问来源于stack exchange,提问作者pfc
相关产品推荐
相关产品推荐

