从Azure Blob加载H5模型权重时出现层数量不匹配错误求助
解决Azure Blob下载模型权重后加载报错的问题
问题分析
你遇到的ValueError: Layer count mismatch when loading weights from file. Model expected 13 layers, found 0 saved layers错误,本质是从Azure Blob下载到本地的临时权重文件是空的或者损坏了——毕竟本地原始文件能正常加载,说明问题出在文件下载/写入环节。
解决方案
1. 替换一次性下载为分块流式写入
readall()在处理较大的权重文件时,可能会因为内存缓冲或网络波动导致文件写入不完整。改用分块读取写入的方式更可靠:
def download_weights_to_temp_file(): """ Downloads the model weights from Azure Blob Storage to a temporary file. Returns the local path to the temporary file. """ try: # Authenticate with Azure default_credential = DefaultAzureCredential() blob_service_client = BlobServiceClient(account_url, credential=default_credential) blob_client = blob_service_client.get_blob_client(container=CONTAINER_NAME_MODEL, blob=MODEL_BLOB_NAME) print(f"Downloading model weights from: {blob_client.url}") # 获取Blob文件大小,用于后续校验 blob_properties = blob_client.get_blob_properties() expected_size = blob_properties.size # Create a temporary file with tempfile.NamedTemporaryFile(delete=False, suffix=".h5") as temp_file: # 分块下载写入,每次读取4MB for chunk in blob_client.download_blob().chunks(chunk_size=4*1024*1024): temp_file.write(chunk) temp_model_path = temp_file.name # 校验文件大小是否匹配 downloaded_size = os.path.getsize(temp_model_path) if downloaded_size != expected_size: raise Exception(f"File size mismatch: expected {expected_size} bytes, got {downloaded_size} bytes") print(f"Model weights downloaded to: {temp_model_path} (size: {downloaded_size} bytes)") return temp_model_path except Exception as e: print(f"Error downloading model weights: {e}") raise
2. 强制写入磁盘避免缓冲遗漏
虽然with语句会自动关闭文件,但某些情况下系统缓冲可能还没完成写入。可以在下载完成后,添加强制写入操作:
with tempfile.NamedTemporaryFile(delete=False, suffix=".h5") as temp_file: for chunk in blob_client.download_blob().chunks(chunk_size=4*1024*1024): temp_file.write(chunk) # 强制把缓冲内容写入磁盘 temp_file.flush() os.fsync(temp_file.fileno()) temp_model_path = temp_file.name
3. 确认Azure Blob文件完整性
检查你上传到Azure Blob的权重文件是否和本地正常文件一致:
- 对比本地文件和Blob文件的MD5哈希值(Azure Blob属性里可以看到Content-MD5)
- 如果哈希不匹配,重新上传本地的原始权重文件到Blob容器
4. 加载前添加空文件检查
在加载权重前,临时检查文件大小,避免加载空文件:
def load_model(): temp_model_path = download_weights_to_temp_file() try: # 额外检查文件大小,避免空文件 if os.path.getsize(temp_model_path) == 0: raise Exception("Downloaded weights file is empty") # Create the model architecture model= create_enet_model(input_shape=(256, 256, 1), num_classes=1) model.compile(optimizer="adam", loss=weighted_binary_crossentropy, metrics=["accuracy"]) model.load_weights(temp_model_path) print(f"Successfully loaded model weights from: {temp_model_path}") except Exception as e: print(f"Error loading model weights: {e}") raise finally: # Clean up the temporary file if os.path.exists(temp_model_path): os.remove(temp_model_path) return model
关键排查点
- 优先检查下载后的文件大小是否和Azure Blob上的一致,这是最快确认文件是否损坏的方法
- 分块下载是解决大文件网络传输问题的通用方案,能避免一次性读取内存溢出或写入不完整
内容的提问来源于stack exchange,提问作者Amanda Spolti
相关产品推荐
相关产品推荐

