| 27 | |
| 28 | template <typename Shard> |
| 29 | bool init_layer_split_runtime(const LayerSplitRuntimeInit & cfg, |
| 30 | std::vector<Shard> & shards, |
| 31 | std::vector<ggml_backend_t> & snapshot_backends) { |
| 32 | const char * log_prefix = cfg.log_prefix ? cfg.log_prefix : "target-split"; |
| 33 | if (!cfg.target_path || !cfg.device || |
| 34 | cfg.device->layer_split_gpus.size() < 2) { |
| 35 | std::fprintf(stderr, "[%s] invalid layer-split config\n", log_prefix); |
| 36 | return false; |
| 37 | } |
| 38 | |
| 39 | const auto info = inspect_gguf_model_info(cfg.target_path); |
| 40 | const int n_layer = info.n_layer; |
| 41 | if (n_layer <= 0) { |
| 42 | std::fprintf(stderr, "[%s] failed to inspect target layer count\n", |
| 43 | log_prefix); |
| 44 | return false; |
| 45 | } |
| 46 | |
| 47 | const auto ranges = compute_layer_ranges( |
| 48 | n_layer, |
| 49 | (int)cfg.device->layer_split_gpus.size(), |
| 50 | cfg.device->layer_split_weights); |
| 51 | if (ranges.size() != cfg.device->layer_split_gpus.size()) { |
| 52 | std::fprintf(stderr, |
| 53 | "[%s] bad layer split for %zu GPUs and %d layers\n", |
| 54 | log_prefix, cfg.device->layer_split_gpus.size(), n_layer); |
| 55 | return false; |
| 56 | } |
| 57 | |
| 58 | shards.resize(cfg.device->layer_split_gpus.size()); |
| 59 | auto shard_metas = layer_split_shard_metas(shards); |
| 60 | if (!init_layer_split_shard_metas( |
| 61 | shard_metas, cfg.device->layer_split_gpus, ranges, |
| 62 | log_prefix)) { |
| 63 | return false; |
| 64 | } |
| 65 | for (size_t i = 0; i < shard_metas.size(); ++i) { |
| 66 | shard_metas[i]->placement_backend = cfg.device->layer_split_backend(i); |
| 67 | } |
| 68 | |
| 69 | (void)enable_layer_split_peer_access( |
| 70 | cfg.device->layer_split_gpus, cfg.device->peer_access); |
| 71 | |
| 72 | return init_layer_split_snapshot_backends( |
| 73 | shard_metas, snapshot_backends, log_prefix); |
| 74 | } |
| 75 | |
| 76 | using LayerSplitForwardStep = std::function<bool( |
| 77 | const std::vector<int32_t> & tokens, |
no test coverage detected