mirror of
https://github.com/leejet/stable-diffusion.cpp.git
synced 2026-10-03 11:27:43 -05:00
feat: add LLaDA-Image support (#1968)
Co-authored-by: leejet <leejet714@gmail.com>
This commit is contained in:
@@ -14,6 +14,7 @@
|
||||
#include "core/util.h"
|
||||
#include "model/diffusion/model.hpp"
|
||||
#include "model/te/clip.hpp"
|
||||
#include "model/te/llada_image_te.hpp"
|
||||
#include "model/te/llm.hpp"
|
||||
#include "model/te/t5.hpp"
|
||||
#include "model_loader.h"
|
||||
@@ -3236,6 +3237,214 @@ struct LTXAVTextProjectionRunner : public GGMLRunner {
|
||||
}
|
||||
};
|
||||
|
||||
// LLaDA-Image's text path is a three-stage pipeline rather than a single encoder pass:
|
||||
// the token embeddings feed a QueryFormer whose 256 queries are appended to the backbone
|
||||
// input, and the backbone's final hidden states are projected to the denoiser's caption dim.
|
||||
// Ref: LLaDAImagePipeline._encode_text.
|
||||
struct LLaDAImageEmbedder : public Conditioner {
|
||||
std::shared_ptr<Tokenizer> tokenizer;
|
||||
std::shared_ptr<LLM::LLMRunner> llm;
|
||||
std::shared_ptr<LLaDAImageTE::QueryFormerRunner> query_former;
|
||||
std::shared_ptr<LLaDAImageTE::TextProjectionRunner> text_projection;
|
||||
std::shared_ptr<LLaDAImageTE::SigVQRunner> sigvq;
|
||||
|
||||
std::string llm_prefix;
|
||||
std::string query_former_prefix;
|
||||
std::string text_projection_prefix;
|
||||
std::string sigvq_prefix;
|
||||
|
||||
LLaDAImageEmbedder(ggml_backend_t backend,
|
||||
const String2TensorStorage& tensor_storage_map = {},
|
||||
const std::string& llm_prefix = "text_encoders.llm",
|
||||
const std::string& query_former_prefix = "queryformer",
|
||||
const std::string& text_projection_prefix = "text_projection",
|
||||
const std::string& sigvq_prefix = "sigvq",
|
||||
std::shared_ptr<RunnerWeightManager> weight_manager = nullptr,
|
||||
const TokenizerConfig& tokenizers = {})
|
||||
: llm_prefix(llm_prefix),
|
||||
query_former_prefix(query_former_prefix),
|
||||
text_projection_prefix(text_projection_prefix),
|
||||
sigvq_prefix(sigvq_prefix) {
|
||||
if (!tokenizers.has(TokenizerConfig::MAIN)) {
|
||||
throw std::runtime_error("LLaDA-Image requires an external LLaDA2 tokenizer.json; pass --tokenizer FILE or set sd_ctx_params_t::tokenizer");
|
||||
}
|
||||
llm = std::make_shared<LLM::LLMRunner>(LLM::LLMArch::LLADA2_MOE,
|
||||
backend,
|
||||
tensor_storage_map,
|
||||
llm_prefix,
|
||||
false,
|
||||
weight_manager);
|
||||
// <|endoftext|> doubles as the pad token in LLaDA2's tokenizer.json.
|
||||
tokenizer = tokenizers.create(TokenizerConfig::MAIN, llm->config.vocab_size, 156892);
|
||||
query_former = std::make_shared<LLaDAImageTE::QueryFormerRunner>(backend,
|
||||
tensor_storage_map,
|
||||
query_former_prefix,
|
||||
weight_manager);
|
||||
text_projection = std::make_shared<LLaDAImageTE::TextProjectionRunner>(backend,
|
||||
tensor_storage_map,
|
||||
text_projection_prefix,
|
||||
weight_manager);
|
||||
|
||||
// SigVQ is only present when the user supplies the editing weights.
|
||||
for (const auto& [name, _] : tensor_storage_map) {
|
||||
if (starts_with(name, sigvq_prefix + ".")) {
|
||||
sigvq = std::make_shared<LLaDAImageTE::SigVQRunner>(backend,
|
||||
tensor_storage_map,
|
||||
sigvq_prefix,
|
||||
weight_manager);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void get_param_tensors(std::map<std::string, ggml_tensor*>& tensors) override {
|
||||
llm->get_param_tensors(tensors, llm_prefix);
|
||||
query_former->get_param_tensors(tensors, query_former_prefix);
|
||||
text_projection->get_param_tensors(tensors, text_projection_prefix);
|
||||
if (sigvq != nullptr) {
|
||||
sigvq->get_param_tensors(tensors, sigvq_prefix);
|
||||
}
|
||||
}
|
||||
|
||||
void get_param_tensor_ops(std::map<ggml_tensor*, enum ggml_op>& tensor_ops) override {
|
||||
llm->get_param_tensor_ops(tensor_ops);
|
||||
}
|
||||
|
||||
void set_flash_attention_enabled(bool enabled) override {
|
||||
llm->set_flash_attention_enabled(enabled);
|
||||
query_former->set_flash_attention_enabled(enabled);
|
||||
text_projection->set_flash_attention_enabled(enabled);
|
||||
if (sigvq != nullptr) {
|
||||
sigvq->set_flash_attention_enabled(enabled);
|
||||
}
|
||||
}
|
||||
|
||||
void set_max_graph_vram_bytes(size_t max_vram_bytes) override {
|
||||
llm->set_max_graph_vram_bytes(max_vram_bytes);
|
||||
query_former->set_max_graph_vram_bytes(max_vram_bytes);
|
||||
text_projection->set_max_graph_vram_bytes(max_vram_bytes);
|
||||
if (sigvq != nullptr) {
|
||||
sigvq->set_max_graph_vram_bytes(max_vram_bytes);
|
||||
}
|
||||
}
|
||||
|
||||
void set_runtime_backends(const std::vector<ggml_backend_t>& backends) override {
|
||||
llm->set_runtime_backends(backends);
|
||||
}
|
||||
|
||||
void set_graph_cut_layer_split_enabled(bool enabled) override {
|
||||
llm->set_graph_cut_layer_split_enabled(enabled);
|
||||
}
|
||||
|
||||
void set_graph_cut_layer_split_backend_vram_limits(const std::vector<size_t>& limits) override {
|
||||
llm->set_graph_cut_layer_split_backend_vram_limits(limits);
|
||||
}
|
||||
|
||||
void get_layer_split_param_tensors(std::map<std::string, ggml_tensor*>& tensors) override {
|
||||
llm->get_param_tensors(tensors, llm_prefix);
|
||||
}
|
||||
|
||||
void set_weight_adapter(const std::shared_ptr<WeightAdapter>& adapter) override {
|
||||
llm->set_weight_adapter(adapter);
|
||||
query_former->set_weight_adapter(adapter);
|
||||
text_projection->set_weight_adapter(adapter);
|
||||
if (sigvq != nullptr) {
|
||||
sigvq->set_weight_adapter(adapter);
|
||||
}
|
||||
}
|
||||
|
||||
void runner_end() override {
|
||||
llm->runner_end();
|
||||
query_former->runner_end();
|
||||
text_projection->runner_end();
|
||||
if (sigvq != nullptr) {
|
||||
sigvq->runner_end();
|
||||
}
|
||||
}
|
||||
|
||||
SDCondition get_learned_condition(int n_threads,
|
||||
const ConditionerParams& conditioner_params) override {
|
||||
const int64_t num_queries = 256;
|
||||
const bool has_ref_images = conditioner_params.ref_images != nullptr && !conditioner_params.ref_images->empty();
|
||||
if (has_ref_images && sigvq == nullptr) {
|
||||
LOG_ERROR("LLaDA-Image editing requires connectors with SigVQ weights");
|
||||
return {};
|
||||
}
|
||||
|
||||
std::string text = conditioner_params.text;
|
||||
while (!text.empty() && std::isspace(static_cast<unsigned char>(text.front()))) {
|
||||
text.erase(text.begin());
|
||||
}
|
||||
while (!text.empty() && std::isspace(static_cast<unsigned char>(text.back()))) {
|
||||
text.pop_back();
|
||||
}
|
||||
std::string prompt = text.empty()
|
||||
? "<role>HUMAN</role> Generate an image.\n<role>ASSISTANT</role>\n<IMAGE1>"
|
||||
: "<role>HUMAN</role> Generate an image: " + text + "\n<role>ASSISTANT</role>\n<IMAGE1>";
|
||||
|
||||
std::vector<int> tokens;
|
||||
if (!tokenizer->encode(prompt, tokens, nullptr)) {
|
||||
return {};
|
||||
}
|
||||
int64_t n_text = static_cast<int64_t>(tokens.size());
|
||||
GGML_ASSERT(n_text > 0);
|
||||
|
||||
sd::Tensor<int32_t> text_ids({n_text}, std::vector<int32_t>(tokens.begin(), tokens.end()));
|
||||
auto inputs_embeds = llm->compute_input_embeds(n_threads, text_ids);
|
||||
auto query_embeds = query_former->compute(n_threads, inputs_embeds);
|
||||
|
||||
// splice_image_embeds() replaces tokens in place, so the query slots have to exist in
|
||||
// input_ids; their ids are irrelevant because the embeddings are overwritten.
|
||||
std::vector<int32_t> padded(tokens.begin(), tokens.end());
|
||||
padded.resize(static_cast<size_t>(n_text + num_queries), tokenizer->PAD_TOKEN_ID);
|
||||
int64_t n_total = static_cast<int64_t>(padded.size());
|
||||
sd::Tensor<int32_t> input_ids({n_total}, padded);
|
||||
|
||||
// Bidirectional everywhere except that the text tokens must not see the appended
|
||||
// queries, matching backbone_attention_mask[:, :, :text_length, text_length:] = min.
|
||||
const float mask_min = std::numeric_limits<float>::lowest() / 4.0f;
|
||||
sd::Tensor<float> attention_mask({n_total, n_total});
|
||||
for (int64_t i1 = 0; i1 < n_total; ++i1) {
|
||||
for (int64_t i0 = 0; i0 < n_total; ++i0) {
|
||||
float value = (i1 < n_text && i0 >= n_text) ? mask_min : 0.0f;
|
||||
attention_mask[i0 + n_total * i1] = value;
|
||||
}
|
||||
}
|
||||
|
||||
LLM::ImageEmbeds image_embeds;
|
||||
image_embeds.emplace_back(static_cast<int>(n_text), query_embeds);
|
||||
|
||||
std::set<int> out_layers = {static_cast<int>(llm->config.num_layers) + 1};
|
||||
auto hidden_states = llm->compute(n_threads,
|
||||
input_ids,
|
||||
attention_mask,
|
||||
image_embeds,
|
||||
out_layers);
|
||||
|
||||
SDCondition result;
|
||||
result.c_crossattn = text_projection->compute(n_threads, hidden_states);
|
||||
|
||||
// Editing: SigVQ sees the reference at half the output resolution, as in
|
||||
// LLaDAImagePipeline._encode_source_image.
|
||||
if (has_ref_images) {
|
||||
const auto& ref = conditioner_params.ref_images->front();
|
||||
auto resized = sd::ops::interpolate(ref,
|
||||
{conditioner_params.width / 2,
|
||||
conditioner_params.height / 2,
|
||||
ref.shape()[2],
|
||||
ref.shape()[3]},
|
||||
sd::ops::InterpolateMode::Bilinear);
|
||||
resized = resized * 2.f - 1.f;
|
||||
auto semantic = sigvq->compute(n_threads, resized);
|
||||
if (semantic.empty()) {
|
||||
return {};
|
||||
}
|
||||
result.extra_c_crossattns.push_back(std::move(semantic));
|
||||
}
|
||||
return result;
|
||||
}
|
||||
};
|
||||
|
||||
struct LTXAVEmbedder : public Conditioner {
|
||||
static constexpr int64_t kHiddenSize = 3840;
|
||||
static constexpr int64_t kNumStates = 49;
|
||||
|
||||
Reference in New Issue
Block a user