mirror of
https://github.com/leejet/stable-diffusion.cpp.git
synced 2026-09-29 01:18:05 -05:00
fix: map mmapped weights through Metal buffers instead of CPU buffers (#2037)
This commit is contained in:
+58
-6
@@ -874,7 +874,8 @@ void ModelLoader::process_model_files(bool enable_mmap, bool writable_mmap) {
|
||||
|
||||
std::vector<MmapTensorStore> ModelLoader::mmap_tensors(std::map<std::string, ggml_tensor*>& tensors,
|
||||
std::set<std::string> ignore_tensors,
|
||||
bool writable_mmap) {
|
||||
bool writable_mmap,
|
||||
ggml_backend_dev_t device) {
|
||||
std::set<std::string> names;
|
||||
for (const auto& entry : tensors) {
|
||||
names.insert(entry.first);
|
||||
@@ -896,6 +897,39 @@ std::vector<MmapTensorStore> ModelLoader::mmap_tensors(std::map<std::string, ggm
|
||||
if (!fdata.mmbuffer)
|
||||
continue;
|
||||
|
||||
// Wrapped on first use: a device buffer makes the whole file resident on that device.
|
||||
std::shared_ptr<struct ggml_backend_buffer> file_buffer = device == nullptr ? fdata.mmbuffer : nullptr;
|
||||
bool file_unmappable = false;
|
||||
|
||||
auto buffer_for_file = [&]() -> ggml_backend_buffer_t {
|
||||
if (file_buffer || file_unmappable) {
|
||||
return file_buffer.get();
|
||||
}
|
||||
auto cached = fdata.device_mmbuffers.find(device);
|
||||
if (cached != fdata.device_mmbuffers.end()) {
|
||||
file_buffer = cached->second;
|
||||
return file_buffer.get();
|
||||
}
|
||||
size_t max_tensor_size = 0;
|
||||
for (const auto& ts : fdata.tensors) {
|
||||
max_tensor_size = std::max(max_tensor_size, static_cast<size_t>(ts.nbytes()));
|
||||
}
|
||||
ggml_backend_buffer_t buf = sd_backend_dev_buffer_from_host_ptr(device,
|
||||
fdata.mmapped->writable_data(),
|
||||
fdata.mmapped->size(),
|
||||
max_tensor_size);
|
||||
if (buf == nullptr) {
|
||||
LOG_WARN("mmap: %s cannot map '%s', loading it instead",
|
||||
ggml_backend_dev_name(device), fdata.path.c_str());
|
||||
file_unmappable = true;
|
||||
return nullptr;
|
||||
}
|
||||
LOG_INFO("mmap: mapped '%s' for %s", fdata.path.c_str(), ggml_backend_dev_name(device));
|
||||
file_buffer = std::shared_ptr<struct ggml_backend_buffer>(buf, ggml_backend_buffer_free);
|
||||
fdata.device_mmbuffers[device] = file_buffer;
|
||||
return file_buffer.get();
|
||||
};
|
||||
|
||||
const std::vector<TensorStorage>& file_tensors = fdata.tensors;
|
||||
|
||||
size_t file_mapped_bytes = 0;
|
||||
@@ -944,10 +978,13 @@ std::vector<MmapTensorStore> ModelLoader::mmap_tensors(std::map<std::string, ggm
|
||||
continue;
|
||||
}
|
||||
|
||||
ggml_backend_buffer_t buf_mmap = fdata.mmbuffer.get();
|
||||
uint8_t* mmap_data = static_cast<uint8_t*>(ggml_backend_buffer_get_base(buf_mmap));
|
||||
dst_tensor->buffer = buf_mmap;
|
||||
dst_tensor->data = mmap_data + tensor_offset;
|
||||
ggml_backend_buffer_t buf_mmap = buffer_for_file();
|
||||
if (buf_mmap == nullptr) {
|
||||
break;
|
||||
}
|
||||
uint8_t* mmap_data = static_cast<uint8_t*>(ggml_backend_buffer_get_base(buf_mmap));
|
||||
dst_tensor->buffer = buf_mmap;
|
||||
dst_tensor->data = mmap_data + tensor_offset;
|
||||
|
||||
file_mapped_bytes += tensor_size;
|
||||
file_mapped_tensors++;
|
||||
@@ -956,7 +993,7 @@ std::vector<MmapTensorStore> ModelLoader::mmap_tensors(std::map<std::string, ggm
|
||||
if (file_mapped_bytes > 0) {
|
||||
mapped_tensors += file_mapped_tensors;
|
||||
mapped_bytes += file_mapped_bytes;
|
||||
result.push_back({fdata.mmapped, fdata.mmbuffer});
|
||||
result.push_back({fdata.mmapped, file_buffer});
|
||||
}
|
||||
}
|
||||
|
||||
@@ -972,6 +1009,16 @@ std::vector<MmapTensorStore> ModelLoader::mmap_tensors(std::map<std::string, ggm
|
||||
return result;
|
||||
}
|
||||
|
||||
std::vector<ggml_backend_buffer_t> ModelLoader::get_device_mmap_buffers() const {
|
||||
std::vector<ggml_backend_buffer_t> buffers;
|
||||
for (const auto& fdata : file_data) {
|
||||
for (const auto& entry : fdata.device_mmbuffers) {
|
||||
buffers.push_back(entry.second.get());
|
||||
}
|
||||
}
|
||||
return buffers;
|
||||
}
|
||||
|
||||
bool ModelLoader::load_tensors(on_new_tensor_cb_t on_new_tensor_cb,
|
||||
bool enable_mmap,
|
||||
const std::set<std::string>* target_tensor_names,
|
||||
@@ -1115,6 +1162,11 @@ bool ModelLoader::load_tensors(on_new_tensor_cb_t on_new_tensor_cb,
|
||||
if (dst_tensor->buffer != nullptr && dst_tensor->buffer == fdata.mmbuffer.get()) {
|
||||
continue;
|
||||
}
|
||||
if (dst_tensor->buffer != nullptr &&
|
||||
std::any_of(fdata.device_mmbuffers.begin(), fdata.device_mmbuffers.end(),
|
||||
[&](const auto& entry) { return entry.second.get() == dst_tensor->buffer; })) {
|
||||
continue;
|
||||
}
|
||||
|
||||
size_t nbytes_to_read = tensor_storage.nbytes_to_read();
|
||||
|
||||
|
||||
Reference in New Issue
Block a user