diff --git a/DESCRIPTION b/DESCRIPTION index 9a28bea..1da70f9 100644 --- a/DESCRIPTION +++ b/DESCRIPTION @@ -11,7 +11,7 @@ Description: A collection of minimal implementations of deep learning License: MIT + file LICENSE Encoding: UTF-8 Roxygen: list(markdown = TRUE) -RoxygenNote: 7.3.2 +RoxygenNote: 7.3.3 Suggests: testthat (>= 3.0.0) Depends: diff --git a/NAMESPACE b/NAMESPACE index c66122d..7cefd74 100644 --- a/NAMESPACE +++ b/NAMESPACE @@ -16,6 +16,12 @@ export(hf_state_dict) export(llama) export(llama_from_config) export(llama_from_pretrained) +export(ministral) +export(ministral_from_config) +export(ministral_from_pretrained) +export(qwen3) +export(qwen3_from_config) +export(qwen3_from_pretrained) import(torch) importFrom(dotty,.) importFrom(hfhub,WEIGHTS_INDEX_NAME) diff --git a/R/ministral.R b/R/ministral.R new file mode 100644 index 0000000..546dc76 --- /dev/null +++ b/R/ministral.R @@ -0,0 +1,331 @@ +# References: +# - https://huggingface.co/mistralai/Ministral-3-14B-Instruct-2512 +# - https://github.com/huggingface/transformers/blob/main/src/transformers/models/mistral/modeling_mistral.py + +#' @noRd +#' @importFrom zeallot %<-% +#' @importFrom purrr map +#' @import torch +NULL + +# YaRN RoPE helper functions +yarn_find_correction_dim <- function(num_rotations, dim, base, max_position_embeddings) { + (dim * log(max_position_embeddings / (num_rotations * 2 * pi))) / (2 * log(base)) +} + +yarn_find_correction_range <- function(low_rot, high_rot, dim, base, max_position_embeddings) { + low <- floor(yarn_find_correction_dim(low_rot, dim, base, max_position_embeddings)) + high <- ceiling(yarn_find_correction_dim(high_rot, dim, base, max_position_embeddings)) + c(max(low, 0), min(high, dim - 1)) +} + +yarn_linear_ramp_mask <- function(min_val, max_val, dim, dtype = torch_float32()) { + if (min_val == max_val) min_val <- min_val - 0.001 + linear_func <- (torch_arange(0, dim - 1, dtype = dtype) - min_val) / (max_val - min_val) + torch_clamp(linear_func, 0, 1) +} + +ministral_rotate_half <- function(x) { + c(x1, x2) %<-% torch_split(x, x$size(-1) / 2, -1) + torch_cat(list(-x2, x1), dim = -1) +} + +repeat_kv <- function(hidden_states, n_rep) { + if (n_rep == 1) return(hidden_states) + c(batch, num_kv_heads, seq_len, head_dim) %<-% hidden_states$shape + hidden_states$unsqueeze(3)$ + expand(c(batch, num_kv_heads, n_rep, seq_len, head_dim))$ + reshape(c(batch, num_kv_heads * n_rep, seq_len, head_dim)) +} + +nn_ministral_rmsnorm <- nn_module( + initialize = function(hidden_size, eps = 1e-5) { + self$weight <- nn_parameter(torch_ones(hidden_size)) + self$eps <- eps + }, + forward = function(x) { + dtype <- x$dtype + variance <- x$to(dtype = "float32")$pow(2)$mean(-1, keepdim = TRUE) + x <- x * torch_rsqrt(variance + self$eps) + (self$weight * x)$to(dtype = dtype) + } +) + +nn_ministral_yarn_rotary_embedding <- nn_module( + initialize = function(head_dim, max_pos, base, factor, beta_fast, beta_slow, + original_max_pos, mscale, mscale_all_dim) { + self$head_dim <- head_dim + self$max_pos <- max_pos + self$base <- base + self$factor <- factor + self$beta_fast <- beta_fast + self$beta_slow <- beta_slow + self$original_max_pos <- original_max_pos + self$mscale <- mscale + self$mscale_all_dim <- mscale_all_dim + self$cached_embeddings() + }, + .load_from_state_dict = function(...) { + super$.load_from_state_dict(...) + self$cached_embeddings(invalidate = TRUE) + }, + get_mscale = function(scale, mscale) { + if (mscale <= 0) return(1.0) + 0.1 * mscale * log(scale) + 1.0 + }, + cached_embeddings = function(t = 1, invalidate = FALSE) { + invalidate <- invalidate || is.null(self$cos) + if (invalidate) { + dim <- self$head_dim + pos_freqs <- self$base ^ (torch_arange(0, dim - 1, step = 2) / dim) + inv_freq_extrapolation <- 1.0 / pos_freqs + inv_freq_interpolation <- 1.0 / (self$factor * pos_freqs) + + c(low, high) %<-% yarn_find_correction_range( + self$beta_slow, self$beta_fast, dim, self$base, self$original_max_pos + ) + inv_freq_extrapolation_factor <- yarn_linear_ramp_mask(low, high, dim / 2) + + inv_freq <- inv_freq_interpolation * (1 - inv_freq_extrapolation_factor) + + inv_freq_extrapolation * inv_freq_extrapolation_factor + self$inv_freq <- nn_buffer(inv_freq, persistent = FALSE) + + self$attention_scale <- self$get_mscale(self$factor, self$mscale) / + self$get_mscale(self$factor, self$mscale_all_dim) + + freqs <- torch_arange(start = 0, end = self$max_pos - 1)$ + float()$outer(self$inv_freq)$view(c(1, 1, self$max_pos, dim / 2)) + emb <- torch_cat(list(freqs, freqs), dim = -1) + self$cos <- nn_buffer(emb$cos(), persistent = FALSE) + self$sin <- nn_buffer(emb$sin(), persistent = FALSE) + } + list(self$cos[,,1:t,], self$sin[,,1:t,], self$attention_scale) + }, + forward = function(x) { + c(b, nh, t, ed) %<-% x$shape + c(cos, sin, attn_scale) %<-% self$cached_embeddings(t) + (x * cos + ministral_rotate_half(x) * sin) * attn_scale + } +) + +nn_ministral_attention <- nn_module( + initialize = function(n_embd, n_head, n_kv_head, head_dim, max_pos, + rope_base, rope_factor, rope_beta_fast, rope_beta_slow, + rope_original_max_pos, rope_mscale, rope_mscale_all_dim) { + self$n_head <- n_head + self$n_kv_head <- n_kv_head + self$head_dim <- head_dim + self$n_kv_groups <- n_head %/% n_kv_head + self$max_pos <- max_pos + + self$rotary <- nn_ministral_yarn_rotary_embedding( + head_dim, max_pos, rope_base, rope_factor, rope_beta_fast, rope_beta_slow, + rope_original_max_pos, rope_mscale, rope_mscale_all_dim + ) + + self$q_proj <- nn_linear(n_embd, n_head * head_dim, bias = FALSE) + self$k_proj <- nn_linear(n_embd, n_kv_head * head_dim, bias = FALSE) + self$v_proj <- nn_linear(n_embd, n_kv_head * head_dim, bias = FALSE) + self$o_proj <- nn_linear(n_head * head_dim, n_embd, bias = FALSE) + self$cached_bias() + }, + forward = function(x) { + c(b, t, h) %<-% x$shape + + q <- self$q_proj(x)$view(c(b, t, self$n_head, self$head_dim))$transpose(2, 3) + k <- self$k_proj(x)$view(c(b, t, self$n_kv_head, self$head_dim))$transpose(2, 3) + v <- self$v_proj(x)$view(c(b, t, self$n_kv_head, self$head_dim))$transpose(2, 3) + + q <- self$rotary(q)$to(dtype = "float") + k <- self$rotary(k)$to(dtype = "float") + + k <- repeat_kv(k, self$n_kv_groups) + v <- repeat_kv(v, self$n_kv_groups) + + att <- torch_matmul(q, k$transpose(-2, -1)) * (1 / sqrt(self$head_dim)) + att <- att$masked_fill(self$bias[,,1:t, 1:t] == 0, self$masked_bias) + att <- nnf_softmax(att, dim = -1)$to(dtype = v$dtype) + + y <- torch_matmul(att, v)$transpose(2, 3)$contiguous()$ + view(c(b, t, self$n_head * self$head_dim)) + self$o_proj(y) + }, + .load_from_state_dict = function(...) { + super$.load_from_state_dict(...) + self$cached_bias() + }, + cached_bias = function() { + self$bias <- torch_ones(self$max_pos, self$max_pos)$bool()$tril()$ + view(c(1, 1, self$max_pos, self$max_pos)) |> nn_buffer(persistent = FALSE) + self$masked_bias <- nn_buffer(torch_scalar_tensor(-Inf), persistent = FALSE) + } +) + +nn_ministral_mlp <- nn_module( + initialize = function(n_embd, n_inter) { + self$gate_proj <- nn_linear(n_embd, n_inter, bias = FALSE) + self$down_proj <- nn_linear(n_inter, n_embd, bias = FALSE) + self$up_proj <- nn_linear(n_embd, n_inter, bias = FALSE) + self$act <- nn_silu() + }, + forward = function(x) { + self$down_proj(self$act(self$gate_proj(x)) * self$up_proj(x)) + } +) + +nn_ministral_layer <- nn_module( + initialize = function(n_embd, n_inter, n_head, n_kv_head, head_dim, max_pos, + rmsnorm_eps, rope_base, rope_factor, rope_beta_fast, + rope_beta_slow, rope_original_max_pos, rope_mscale, + rope_mscale_all_dim) { + self$ln_1 <- nn_ministral_rmsnorm(n_embd, rmsnorm_eps) + self$ln_2 <- nn_ministral_rmsnorm(n_embd, rmsnorm_eps) + self$attn <- nn_ministral_attention( + n_embd, n_head, n_kv_head, head_dim, max_pos, rope_base, rope_factor, + rope_beta_fast, rope_beta_slow, rope_original_max_pos, rope_mscale, + rope_mscale_all_dim + ) + self$mlp <- nn_ministral_mlp(n_embd, n_inter) + }, + forward = function(x) { + x <- x + self$attn(self$ln_1(x)) + x + self$mlp(self$ln_2(x)) + } +) + +nn_ministral_model <- nn_module( + initialize = function(vocab_size, n_embd, n_inter, n_head, n_kv_head, head_dim, + n_layer, max_pos, rmsnorm_eps, rope_base, rope_factor, + rope_beta_fast, rope_beta_slow, rope_original_max_pos, + rope_mscale, rope_mscale_all_dim) { + self$transformer <- nn_module_dict(list( + wte = nn_embedding(vocab_size, n_embd), + h = nn_sequential(!!!map( + 1:n_layer, + \(x) nn_ministral_layer( + n_embd, n_inter, n_head, n_kv_head, head_dim, max_pos, rmsnorm_eps, + rope_base, rope_factor, rope_beta_fast, rope_beta_slow, + rope_original_max_pos, rope_mscale, rope_mscale_all_dim + ) + )), + ln_f = nn_ministral_rmsnorm(n_embd, rmsnorm_eps) + )) + self$lm_head <- nn_linear(n_embd, vocab_size, bias = FALSE) + }, + forward = function(idx) { + x <- self$transformer$wte(idx) + x <- self$transformer$h(x) + x <- self$transformer$ln_f(x) + self$lm_head(x) + } +) + +#' ministral +#' +#' Initializes a Ministral-like model with YaRN RoPE and GQA +#' +#' @param vocab_size Vocabulary size. +#' @param n_embd Embedding dimension. +#' @param n_inter Intermediate size in MLP. +#' @param n_head Number of attention heads. +#' @param n_kv_head Number of key/value heads (for GQA). +#' @param head_dim Dimension of each attention head. +#' @param n_layer Number of transformer layers. +#' @param max_pos Maximum position embeddings. +#' @param rmsnorm_eps Epsilon for RMSNorm. +#' @param rope_base Base for rotary embeddings. +#' @param rope_factor YaRN scaling factor. +#' @param rope_beta_fast YaRN beta_fast parameter. +#' @param rope_beta_slow YaRN beta_slow parameter. +#' @param rope_original_max_pos Original max position embeddings for YaRN. +#' @param rope_mscale YaRN mscale parameter. +#' @param rope_mscale_all_dim YaRN mscale_all_dim parameter. +#' @param identifier HuggingFace model identifier. +#' @param revision HuggingFace model revision. +#' @returns An initialized [torch::nn_module()]. +#' @export +ministral <- function(vocab_size = 131072, n_embd = 5120, n_inter = 16384, + n_head = 32, n_kv_head = 8, head_dim = 128, n_layer = 40, + max_pos = 262144, rmsnorm_eps = 1e-5, rope_base = 1e9, + rope_factor = 16, rope_beta_fast = 32, rope_beta_slow = 1, + rope_original_max_pos = 16384, rope_mscale = 1, + rope_mscale_all_dim = 1) { + nn_ministral_model( + vocab_size, n_embd, n_inter, n_head, n_kv_head, head_dim, n_layer, max_pos, + rmsnorm_eps, rope_base, rope_factor, rope_beta_fast, rope_beta_slow, + rope_original_max_pos, rope_mscale, rope_mscale_all_dim + ) +} + +#' @describeIn ministral Initializes from HuggingFace config +#' @export +ministral_from_config <- function(identifier, revision = "main") { + path <- hfhub::hub_download(identifier, "config.json", revision = revision) + config <- jsonlite::fromJSON(path) + + # Handle multimodal config (text_config nested) + if (!is.null(config$text_config)) { + rope_params <- config$text_config$rope_parameters + config <- config$text_config + } else { + rope_params <- config$rope_parameters %||% config + } + + if (!config$model_type %in% c("ministral3", "mistral")) + cli::cli_abort("Unsupported model_type: {.val {config$model_type}}") + + if (config$hidden_act != "silu") + cli::cli_abort("Unsupported hidden_act: {.val {config$hidden_act}}") + + ministral( + vocab_size = config$vocab_size, + n_embd = config$hidden_size, + n_inter = config$intermediate_size, + n_head = config$num_attention_heads, + n_kv_head = config$num_key_value_heads, + head_dim = config$head_dim %||% (config$hidden_size %/% config$num_attention_heads), + n_layer = config$num_hidden_layers, + max_pos = config$max_position_embeddings, + rmsnorm_eps = config$rms_norm_eps %||% 1e-5, + rope_base = rope_params$rope_theta %||% 1e9, + rope_factor = rope_params$factor %||% 16, + rope_beta_fast = rope_params$beta_fast %||% 32, + rope_beta_slow = rope_params$beta_slow %||% 1, + rope_original_max_pos = rope_params$original_max_position_embeddings %||% 16384, + rope_mscale = rope_params$mscale %||% 1, + rope_mscale_all_dim = rope_params$mscale_all_dim %||% 1 + ) +} + +#' @describeIn ministral Initializes and loads pretrained weights from HF Hub +#' @export +ministral_from_pretrained <- function(identifier, revision = "main") { + with_device(device = "meta", { + model <- ministral_from_config(identifier, revision) + }) + state_dict <- hf_state_dict(identifier, revision) + state_dict <- ministral_hf_weights_remap(state_dict) + model$load_state_dict(state_dict, .refer_to_state_dict = TRUE) + model +} + +ministral_hf_weights_remap <- function(state_dict) { + nms <- names(state_dict) + + # Handle multimodal models (language_model prefix) + nms <- gsub("^language_model\\.", "", nms) + + # Standard remapping + nms <- gsub("model.embed_tokens.weight", "transformer.wte.weight", nms, fixed = TRUE) + nms <- gsub("model.layers", "transformer.h", nms, fixed = TRUE) + nms <- gsub("self_attn", "attn", nms, fixed = TRUE) + nms <- gsub("input_layernorm", "ln_1", nms, fixed = TRUE) + nms <- gsub("post_attention_layernorm", "ln_2", nms, fixed = TRUE) + nms <- gsub("model.norm", "transformer.ln_f", nms, fixed = TRUE) + + names(state_dict) <- nms + + # Filter out non-language model weights + keep <- !grepl("vision_tower|multi_modal_projector|scale", names(state_dict)) + state_dict[keep] +} diff --git a/R/qwen3.R b/R/qwen3.R new file mode 100644 index 0000000..5103409 --- /dev/null +++ b/R/qwen3.R @@ -0,0 +1,484 @@ +# References: +# - https://huggingface.co/Qwen/Qwen3.5-0.8B +# - https://github.com/huggingface/transformers/blob/main/src/transformers/models/qwen3_5/modeling_qwen3_5.py +# - https://github.com/huggingface/transformers/blob/main/src/transformers/models/qwen3_next/modeling_qwen3_next.py + +#' @noRd +#' @importFrom zeallot %<-% +#' @importFrom purrr map +#' @import torch +NULL + +# Qwen3Next-style RMSNorm with unit offset (weight initialized to zeros, output = (1 + w) * norm(x)) +nn_qwen3_rmsnorm <- nn_module( + initialize = function(hidden_size, eps = 1e-6) { + self$weight <- nn_parameter(torch_zeros(hidden_size)) + self$eps <- eps + }, + forward = function(x) { + dtype <- x$dtype + variance <- x$to(dtype = "float32")$pow(2)$mean(-1, keepdim = TRUE) + x <- x * torch_rsqrt(variance + self$eps) + ((1 + self$weight$float()) * x)$to(dtype = dtype) + } +) + +# Gated RMSNorm: normalizes hidden_states then multiplies by silu(gate) +nn_qwen3_rmsnorm_gated <- nn_module( + initialize = function(hidden_size, eps = 1e-6) { + self$weight <- nn_parameter(torch_ones(hidden_size)) + self$eps <- eps + }, + forward = function(hidden_states, gate) { + dtype <- hidden_states$dtype + hidden_states <- hidden_states$to(dtype = "float32") + variance <- hidden_states$pow(2)$mean(-1, keepdim = TRUE) + hidden_states <- hidden_states * torch_rsqrt(variance + self$eps) + hidden_states <- (self$weight * hidden_states)$to(dtype = dtype) + hidden_states <- hidden_states * nnf_silu(gate$to(dtype = "float32")) + hidden_states$to(dtype = dtype) + } +) + +qwen3_rotate_half <- function(x) { + c(x1, x2) %<-% torch_split(x, x$size(-1) / 2, -1) + torch_cat(list(-x2, x1), dim = -1) +} + +# RoPE with partial rotary factor: only rotates the first `rotary_dim` dims +# Computes cos/sin on the fly to avoid large pre-allocated buffers +nn_qwen3_rotary_embedding <- nn_module( + initialize = function(head_dim, max_pos, base = 10000, partial_rotary_factor = 1.0) { + self$head_dim <- head_dim + self$max_pos <- max_pos + self$base <- base + self$rotary_dim <- as.integer(head_dim * partial_rotary_factor) + + self$compute_inv_freq() + }, + compute_inv_freq = function() { + dim <- self$rotary_dim + inv_freq <- 1 / (self$base ^ (torch_arange(0, dim - 1, step = 2) / dim)) + self$inv_freq <- nn_buffer(inv_freq, persistent = FALSE) + }, + .load_from_state_dict = function(...) { + super$.load_from_state_dict(...) + self$compute_inv_freq() + }, + forward = function(x) { + c(b, nh, t, d) %<-% x$shape + dim <- self$rotary_dim + + freqs <- torch_arange(start = 0, end = t - 1, device = x$device)$ + float()$outer(self$inv_freq)$ + view(c(1, 1, t, dim / 2)) + emb <- torch_cat(list(freqs, freqs), dim = -1) + cos <- emb$cos() + sin <- emb$sin() + + if (self$rotary_dim < self$head_dim) { + x_rot <- x[,,,1:self$rotary_dim] + x_pass <- x[,,,(self$rotary_dim + 1):d] + x_rot <- x_rot * cos + qwen3_rotate_half(x_rot) * sin + torch_cat(list(x_rot, x_pass), dim = -1) + } else { + x * cos + qwen3_rotate_half(x) * sin + } + } +) + +qwen3_repeat_kv <- function(hidden_states, n_rep) { + if (n_rep == 1) return(hidden_states) + c(batch, num_kv_heads, seq_len, head_dim) %<-% hidden_states$shape + hidden_states$unsqueeze(3)$ + expand(c(batch, num_kv_heads, n_rep, seq_len, head_dim))$ + reshape(c(batch, num_kv_heads * n_rep, seq_len, head_dim)) +} + +# Full attention with output gate (Qwen3Next-style) +# q_proj outputs 2x: half for query, half for sigmoid gate on output +nn_qwen3_attention <- nn_module( + initialize = function(n_embd, n_head, n_kv_head, head_dim, max_pos, + rope_base, partial_rotary_factor, rmsnorm_eps) { + self$n_head <- n_head + self$n_kv_head <- n_kv_head + self$head_dim <- head_dim + self$n_kv_groups <- n_head %/% n_kv_head + + self$rotary <- nn_qwen3_rotary_embedding(head_dim, max_pos, rope_base, partial_rotary_factor) + + # q_proj outputs 2x head_dim: first half = query, second half = gate + self$q_proj <- nn_linear(n_embd, n_head * head_dim * 2, bias = FALSE) + self$k_proj <- nn_linear(n_embd, n_kv_head * head_dim, bias = FALSE) + self$v_proj <- nn_linear(n_embd, n_kv_head * head_dim, bias = FALSE) + self$o_proj <- nn_linear(n_head * head_dim, n_embd, bias = FALSE) + + self$q_norm <- nn_qwen3_rmsnorm(head_dim, eps = rmsnorm_eps) + self$k_norm <- nn_qwen3_rmsnorm(head_dim, eps = rmsnorm_eps) + }, + forward = function(x) { + c(b, t, h) %<-% x$shape + + # q_proj -> split into query and gate + qg <- self$q_proj(x)$view(c(b, t, self$n_head, self$head_dim * 2)) + c(q, gate) %<-% torch_chunk(qg, 2, dim = -1) # each (b, t, n_head, head_dim) + gate <- gate$reshape(c(b, t, -1)) # (b, t, n_head * head_dim) + + q <- self$q_norm(q)$transpose(2, 3) # (b, n_head, t, head_dim) + k <- self$k_proj(x)$view(c(b, t, self$n_kv_head, self$head_dim)) + k <- self$k_norm(k)$transpose(2, 3) + v <- self$v_proj(x)$view(c(b, t, self$n_kv_head, self$head_dim))$transpose(2, 3) + + q <- self$rotary(q)$to(dtype = "float") + k <- self$rotary(k)$to(dtype = "float") + + k <- qwen3_repeat_kv(k, self$n_kv_groups) + v <- qwen3_repeat_kv(v, self$n_kv_groups) + + att <- torch_matmul(q, k$transpose(-2, -1)) * (1 / sqrt(self$head_dim)) + # Generate causal mask on the fly (max_pos is too large to pre-allocate) + causal_mask <- torch_ones(t, t, dtype = torch_bool(), device = x$device)$triu(diagonal = 1) + att <- att$masked_fill(causal_mask$view(c(1, 1, t, t)), -Inf) + att <- nnf_softmax(att, dim = -1)$to(dtype = v$dtype) + + y <- torch_matmul(att, v)$transpose(2, 3)$contiguous()$ + reshape(c(b, t, -1)) + + # Apply sigmoid gate + y <- y * torch_sigmoid(gate) + self$o_proj(y) + } +) + +# Recurrent gated delta rule (simple loop-based implementation) +qwen3_recurrent_gated_delta_rule <- function(query, key, value, g, beta) { + # Inputs: (batch, seq, heads, dim) for q/k/v; (batch, seq, heads) for g/beta + # L2 normalize q and k + query <- nnf_normalize(query, p = 2, dim = -1, eps = 1e-6) + key <- nnf_normalize(key, p = 2, dim = -1, eps = 1e-6) + + # Transpose to (batch, heads, seq, dim) + query <- query$transpose(2, 3)$contiguous()$to(dtype = "float32") + key <- key$transpose(2, 3)$contiguous()$to(dtype = "float32") + value <- value$transpose(2, 3)$contiguous()$to(dtype = "float32") + beta <- beta$transpose(2, 3)$contiguous()$to(dtype = "float32") + g <- g$transpose(2, 3)$contiguous()$to(dtype = "float32") + + c(batch_size, num_heads, seq_len, k_head_dim) %<-% key$shape + v_head_dim <- value$shape[4] + + scale <- 1 / sqrt(k_head_dim) + query <- query * scale + + state <- torch_zeros(batch_size, num_heads, k_head_dim, v_head_dim, + device = query$device, dtype = query$dtype) + outputs <- vector("list", seq_len) + + for (i in seq_len(seq_len)) { + q_t <- query[,,i,] # (batch, heads, k_dim) + k_t <- key[,,i,] # (batch, heads, k_dim) + v_t <- value[,,i,] # (batch, heads, v_dim) + g_t <- g[,,i]$exp()$unsqueeze(-1)$unsqueeze(-1) # (batch, heads, 1, 1) + beta_t <- beta[,,i]$unsqueeze(-1) # (batch, heads, 1) + + state <- state * g_t + # kv_mem = (state * k_t[:,:,:,None]).sum(dim=-2) => k_t @ state + kv_mem <- torch_einsum("bhkv,bhk->bhv", list(state, k_t)) + delta <- (v_t - kv_mem) * beta_t + # state += k_t[:,:,:,None] * delta[:,:,None,:] + state <- state + torch_einsum("bhk,bhv->bhkv", list(k_t, delta)) + # output = (state * q_t[:,:,:,None]).sum(dim=-2) => q_t @ state + outputs[[i]] <- torch_einsum("bhkv,bhk->bhv", list(state, q_t)) + } + + # Stack: list of (batch, heads, v_dim) -> (batch, heads, seq, v_dim) + out <- torch_stack(outputs, dim = 3) + # Back to (batch, seq, heads, v_dim) + out$transpose(2, 3)$contiguous() +} + +# Gated Delta Net linear attention +nn_qwen3_gated_delta_net <- nn_module( + initialize = function(n_embd, n_k_heads, n_v_heads, k_head_dim, v_head_dim, + conv_kernel_size, rmsnorm_eps, hidden_act = "silu") { + self$n_embd <- n_embd + self$n_k_heads <- n_k_heads + self$n_v_heads <- n_v_heads + self$k_head_dim <- k_head_dim + self$v_head_dim <- v_head_dim + self$key_dim <- n_k_heads * k_head_dim + self$value_dim <- n_v_heads * v_head_dim + + conv_dim <- self$key_dim * 2 + self$value_dim + self$conv1d <- nn_conv1d( + in_channels = conv_dim, out_channels = conv_dim, + kernel_size = conv_kernel_size, groups = conv_dim, + bias = FALSE, padding = conv_kernel_size - 1 + ) + + self$dt_bias <- nn_parameter(torch_ones(n_v_heads)) + self$A_log <- nn_parameter(torch_log(torch_empty(n_v_heads)$uniform_(0, 16))) + + self$norm <- nn_qwen3_rmsnorm_gated(v_head_dim, eps = rmsnorm_eps) + self$out_proj <- nn_linear(self$value_dim, n_embd, bias = FALSE) + + self$in_proj_qkv <- nn_linear(n_embd, self$key_dim * 2 + self$value_dim, bias = FALSE) + self$in_proj_z <- nn_linear(n_embd, self$value_dim, bias = FALSE) + self$in_proj_b <- nn_linear(n_embd, n_v_heads, bias = FALSE) + self$in_proj_a <- nn_linear(n_embd, n_v_heads, bias = FALSE) + }, + forward = function(x) { + c(batch_size, seq_len, .) %<-% x$shape + + mixed_qkv <- self$in_proj_qkv(x)$transpose(2, 3) # (batch, channels, seq) + z <- self$in_proj_z(x)$reshape(c(batch_size, seq_len, -1, self$v_head_dim)) + b <- self$in_proj_b(x) + a <- self$in_proj_a(x) + + # Causal conv1d + SiLU + mixed_qkv <- nnf_silu(self$conv1d(mixed_qkv)[,,1:seq_len]) + mixed_qkv <- mixed_qkv$transpose(2, 3) # back to (batch, seq, channels) + + c(query, key, value) %<-% torch_split( + mixed_qkv, c(self$key_dim, self$key_dim, self$value_dim), dim = -1 + ) + + query <- query$reshape(c(batch_size, seq_len, -1, self$k_head_dim)) + key <- key$reshape(c(batch_size, seq_len, -1, self$k_head_dim)) + value <- value$reshape(c(batch_size, seq_len, -1, self$v_head_dim)) + + beta <- b$sigmoid() + g <- -self$A_log$float()$exp() * nnf_softplus(a$float() + self$dt_bias) + + # Repeat interleave k heads to match v heads + if (self$n_v_heads %/% self$n_k_heads > 1) { + n_rep <- self$n_v_heads %/% self$n_k_heads + query <- query$repeat_interleave(n_rep, dim = 3) + key <- key$repeat_interleave(n_rep, dim = 3) + } + + initial_dtype <- query$dtype + core_out <- qwen3_recurrent_gated_delta_rule(query, key, value, g, beta) + + # Gated RMSNorm: reshape to 2D, apply, reshape back + core_out_2d <- core_out$reshape(c(-1, self$v_head_dim)) + z_2d <- z$reshape(c(-1, self$v_head_dim)) + core_out_2d <- self$norm(core_out_2d, z_2d) + core_out <- core_out_2d$reshape(c(batch_size, seq_len, -1))$to(dtype = initial_dtype) + + self$out_proj(core_out) + } +) + +nn_qwen3_mlp <- nn_module( + initialize = function(n_embd, n_inter) { + self$gate_proj <- nn_linear(n_embd, n_inter, bias = FALSE) + self$down_proj <- nn_linear(n_inter, n_embd, bias = FALSE) + self$up_proj <- nn_linear(n_embd, n_inter, bias = FALSE) + self$act <- nn_silu() + }, + forward = function(x) { + self$down_proj(self$act(self$gate_proj(x)) * self$up_proj(x)) + } +) + +nn_qwen3_layer <- nn_module( + initialize = function(layer_type, n_embd, n_inter, n_head, n_kv_head, head_dim, + max_pos, rmsnorm_eps, rope_base, partial_rotary_factor, + n_k_heads, n_v_heads, k_head_dim, v_head_dim, + conv_kernel_size) { + self$layer_type <- layer_type + self$ln_1 <- nn_qwen3_rmsnorm(n_embd, rmsnorm_eps) + self$ln_2 <- nn_qwen3_rmsnorm(n_embd, rmsnorm_eps) + + if (layer_type == "full_attention") { + self$attn <- nn_qwen3_attention( + n_embd, n_head, n_kv_head, head_dim, max_pos, + rope_base, partial_rotary_factor, rmsnorm_eps + ) + } else { + self$linear_attn <- nn_qwen3_gated_delta_net( + n_embd, n_k_heads, n_v_heads, k_head_dim, v_head_dim, + conv_kernel_size, rmsnorm_eps + ) + } + + self$mlp <- nn_qwen3_mlp(n_embd, n_inter) + }, + forward = function(x) { + if (self$layer_type == "full_attention") { + x <- x + self$attn(self$ln_1(x)) + } else { + x <- x + self$linear_attn(self$ln_1(x)) + } + x + self$mlp(self$ln_2(x)) + } +) + +nn_qwen3_model <- nn_module( + initialize = function(vocab_size, n_embd, n_inter, n_head, n_kv_head, head_dim, + n_layer, max_pos, rmsnorm_eps, rope_base, + partial_rotary_factor, layer_types, + n_k_heads, n_v_heads, k_head_dim, v_head_dim, + conv_kernel_size) { + self$transformer <- nn_module_dict(list( + wte = nn_embedding(vocab_size, n_embd), + h = nn_sequential(!!!map( + seq_len(n_layer), + \(i) nn_qwen3_layer( + layer_type = layer_types[i], + n_embd = n_embd, n_inter = n_inter, + n_head = n_head, n_kv_head = n_kv_head, head_dim = head_dim, + max_pos = max_pos, rmsnorm_eps = rmsnorm_eps, + rope_base = rope_base, partial_rotary_factor = partial_rotary_factor, + n_k_heads = n_k_heads, n_v_heads = n_v_heads, + k_head_dim = k_head_dim, v_head_dim = v_head_dim, + conv_kernel_size = conv_kernel_size + ) + )), + ln_f = nn_qwen3_rmsnorm(n_embd, rmsnorm_eps) + )) + self$lm_head <- nn_linear(n_embd, vocab_size, bias = FALSE) + }, + forward = function(idx) { + x <- self$transformer$wte(idx) + x <- self$transformer$h(x) + x <- self$transformer$ln_f(x) + self$lm_head(x) + } +) + +#' Qwen3.5 +#' +#' Initializes a Qwen3.5 model with hybrid linear/full attention (Gated Delta Net) +#' +#' @param vocab_size Vocabulary size. +#' @param n_embd Embedding dimension. +#' @param n_inter Intermediate size in MLP. +#' @param n_head Number of attention heads (full attention). +#' @param n_kv_head Number of key/value heads (full attention GQA). +#' @param head_dim Dimension of each attention head (full attention). +#' @param n_layer Number of transformer layers. +#' @param max_pos Maximum position embeddings. +#' @param rmsnorm_eps Epsilon for RMSNorm. +#' @param rope_base Base for rotary embeddings. +#' @param partial_rotary_factor Fraction of head_dim to apply rotary embeddings to. +#' @param layer_types Character vector of layer types ("linear_attention" or "full_attention"). +#' @param n_k_heads Number of key heads (linear attention). +#' @param n_v_heads Number of value heads (linear attention). +#' @param k_head_dim Key head dimension (linear attention). +#' @param v_head_dim Value head dimension (linear attention). +#' @param conv_kernel_size Causal conv1d kernel size (linear attention). +#' @param identifier HuggingFace model identifier. +#' @param revision HuggingFace model revision. +#' @returns An initialized [torch::nn_module()]. +#' @export +qwen3 <- function(vocab_size = 248320, n_embd = 1024, n_inter = 3584, + n_head = 8, n_kv_head = 2, head_dim = 256, + n_layer = 24, max_pos = 262144, rmsnorm_eps = 1e-6, + rope_base = 1e7, partial_rotary_factor = 0.25, + layer_types = rep(c(rep("linear_attention", 3), "full_attention"), 6), + n_k_heads = 16, n_v_heads = 16, + k_head_dim = 128, v_head_dim = 128, + conv_kernel_size = 4) { + nn_qwen3_model( + vocab_size, n_embd, n_inter, n_head, n_kv_head, head_dim, + n_layer, max_pos, rmsnorm_eps, rope_base, partial_rotary_factor, + layer_types, n_k_heads, n_v_heads, k_head_dim, v_head_dim, conv_kernel_size + ) +} + +#' @describeIn qwen3 Initializes from HuggingFace config +#' @export +qwen3_from_config <- function(identifier, revision = "main") { + path <- hfhub::hub_download(identifier, "config.json", revision = revision) + config <- jsonlite::fromJSON(path) + + # Handle multimodal config (text_config nested) + if (!is.null(config$text_config)) { + text_config <- config$text_config + } else { + text_config <- config + } + + if (!text_config$model_type %in% c("qwen3_5_text", "qwen3_5")) + cli::cli_abort("Unsupported model_type: {.val {text_config$model_type}}") + + if (text_config$hidden_act != "silu") + cli::cli_abort("Unsupported hidden_act: {.val {text_config$hidden_act}}") + + rope_params <- text_config$rope_parameters %||% text_config + + # Build layer_types from full_attention_interval + full_attn_interval <- text_config$full_attention_interval %||% 4 + n_layer <- text_config$num_hidden_layers + layer_types <- ifelse( + seq_len(n_layer) %% full_attn_interval == 0, + "full_attention", + "linear_attention" + ) + + # Linear attention config + linear_config <- text_config$linear_attention_config %||% list() + + qwen3( + vocab_size = text_config$vocab_size, + n_embd = text_config$hidden_size, + n_inter = text_config$intermediate_size, + n_head = text_config$num_attention_heads, + n_kv_head = text_config$num_key_value_heads, + head_dim = text_config$head_dim %||% (text_config$hidden_size %/% text_config$num_attention_heads), + n_layer = n_layer, + max_pos = text_config$max_position_embeddings, + rmsnorm_eps = text_config$rms_norm_eps %||% 1e-6, + rope_base = rope_params$rope_theta %||% 1e7, + partial_rotary_factor = rope_params$partial_rotary_factor %||% 0.25, + layer_types = layer_types, + n_k_heads = linear_config$num_key_heads %||% 16, + n_v_heads = linear_config$num_value_heads %||% 16, + k_head_dim = linear_config$key_head_dim %||% 128, + v_head_dim = linear_config$value_head_dim %||% 128, + conv_kernel_size = linear_config$conv_kernel_dim %||% 4 + ) +} + +#' @describeIn qwen3 Initializes and loads pretrained weights from HF Hub +#' @export +qwen3_from_pretrained <- function(identifier, revision = "main") { + with_device(device = "meta", { + model <- qwen3_from_config(identifier, revision) + }) + state_dict <- hf_state_dict(identifier, revision) + state_dict <- qwen3_hf_weights_remap(state_dict) + model$load_state_dict(state_dict, .refer_to_state_dict = TRUE) + model +} + +qwen3_hf_weights_remap <- function(state_dict) { + nms <- names(state_dict) + + # Filter out non-text weights (vision, projector, mtp, etc.) + keep <- !grepl("visual|vision|multi_modal|mtp\\.", nms) + state_dict <- state_dict[keep] + nms <- names(state_dict) + + # Strip multimodal prefix: model.language_model. -> model. + nms <- gsub("^model\\.language_model\\.", "model.", nms) + + # Standard remapping + nms <- gsub("model.embed_tokens.weight", "transformer.wte.weight", nms, fixed = TRUE) + nms <- gsub("model.layers", "transformer.h", nms, fixed = TRUE) + nms <- gsub("self_attn", "attn", nms, fixed = TRUE) + nms <- gsub("input_layernorm", "ln_1", nms, fixed = TRUE) + nms <- gsub("post_attention_layernorm", "ln_2", nms, fixed = TRUE) + nms <- gsub("model.norm", "transformer.ln_f", nms, fixed = TRUE) + + names(state_dict) <- nms + + # Tie lm_head to embedding if not present + if (!"lm_head.weight" %in% names(state_dict)) { + state_dict["lm_head.weight"] <- state_dict["transformer.wte.weight"] + } + + state_dict +} diff --git a/man/ministral.Rd b/man/ministral.Rd new file mode 100644 index 0000000..5b946d3 --- /dev/null +++ b/man/ministral.Rd @@ -0,0 +1,81 @@ +% Generated by roxygen2: do not edit by hand +% Please edit documentation in R/ministral.R +\name{ministral} +\alias{ministral} +\alias{ministral_from_config} +\alias{ministral_from_pretrained} +\title{ministral} +\usage{ +ministral( + vocab_size = 131072, + n_embd = 5120, + n_inter = 16384, + n_head = 32, + n_kv_head = 8, + head_dim = 128, + n_layer = 40, + max_pos = 262144, + rmsnorm_eps = 1e-05, + rope_base = 1e+09, + rope_factor = 16, + rope_beta_fast = 32, + rope_beta_slow = 1, + rope_original_max_pos = 16384, + rope_mscale = 1, + rope_mscale_all_dim = 1 +) + +ministral_from_config(identifier, revision = "main") + +ministral_from_pretrained(identifier, revision = "main") +} +\arguments{ +\item{vocab_size}{Vocabulary size.} + +\item{n_embd}{Embedding dimension.} + +\item{n_inter}{Intermediate size in MLP.} + +\item{n_head}{Number of attention heads.} + +\item{n_kv_head}{Number of key/value heads (for GQA).} + +\item{head_dim}{Dimension of each attention head.} + +\item{n_layer}{Number of transformer layers.} + +\item{max_pos}{Maximum position embeddings.} + +\item{rmsnorm_eps}{Epsilon for RMSNorm.} + +\item{rope_base}{Base for rotary embeddings.} + +\item{rope_factor}{YaRN scaling factor.} + +\item{rope_beta_fast}{YaRN beta_fast parameter.} + +\item{rope_beta_slow}{YaRN beta_slow parameter.} + +\item{rope_original_max_pos}{Original max position embeddings for YaRN.} + +\item{rope_mscale}{YaRN mscale parameter.} + +\item{rope_mscale_all_dim}{YaRN mscale_all_dim parameter.} + +\item{identifier}{HuggingFace model identifier.} + +\item{revision}{HuggingFace model revision.} +} +\value{ +An initialized \code{\link[torch:nn_module]{torch::nn_module()}}. +} +\description{ +Initializes a Ministral-like model with YaRN RoPE and GQA +} +\section{Functions}{ +\itemize{ +\item \code{ministral_from_config()}: Initializes from HuggingFace config + +\item \code{ministral_from_pretrained()}: Initializes and loads pretrained weights from HF Hub + +}} diff --git a/man/qwen3.Rd b/man/qwen3.Rd new file mode 100644 index 0000000..0a5531f --- /dev/null +++ b/man/qwen3.Rd @@ -0,0 +1,84 @@ +% Generated by roxygen2: do not edit by hand +% Please edit documentation in R/qwen3.R +\name{qwen3} +\alias{qwen3} +\alias{qwen3_from_config} +\alias{qwen3_from_pretrained} +\title{Qwen3.5} +\usage{ +qwen3( + vocab_size = 248320, + n_embd = 1024, + n_inter = 3584, + n_head = 8, + n_kv_head = 2, + head_dim = 256, + n_layer = 24, + max_pos = 262144, + rmsnorm_eps = 1e-06, + rope_base = 1e+07, + partial_rotary_factor = 0.25, + layer_types = rep(c(rep("linear_attention", 3), "full_attention"), 6), + n_k_heads = 16, + n_v_heads = 16, + k_head_dim = 128, + v_head_dim = 128, + conv_kernel_size = 4 +) + +qwen3_from_config(identifier, revision = "main") + +qwen3_from_pretrained(identifier, revision = "main") +} +\arguments{ +\item{vocab_size}{Vocabulary size.} + +\item{n_embd}{Embedding dimension.} + +\item{n_inter}{Intermediate size in MLP.} + +\item{n_head}{Number of attention heads (full attention).} + +\item{n_kv_head}{Number of key/value heads (full attention GQA).} + +\item{head_dim}{Dimension of each attention head (full attention).} + +\item{n_layer}{Number of transformer layers.} + +\item{max_pos}{Maximum position embeddings.} + +\item{rmsnorm_eps}{Epsilon for RMSNorm.} + +\item{rope_base}{Base for rotary embeddings.} + +\item{partial_rotary_factor}{Fraction of head_dim to apply rotary embeddings to.} + +\item{layer_types}{Character vector of layer types ("linear_attention" or "full_attention").} + +\item{n_k_heads}{Number of key heads (linear attention).} + +\item{n_v_heads}{Number of value heads (linear attention).} + +\item{k_head_dim}{Key head dimension (linear attention).} + +\item{v_head_dim}{Value head dimension (linear attention).} + +\item{conv_kernel_size}{Causal conv1d kernel size (linear attention).} + +\item{identifier}{HuggingFace model identifier.} + +\item{revision}{HuggingFace model revision.} +} +\value{ +An initialized \code{\link[torch:nn_module]{torch::nn_module()}}. +} +\description{ +Initializes a Qwen3.5 model with hybrid linear/full attention (Gated Delta Net) +} +\section{Functions}{ +\itemize{ +\item \code{qwen3_from_config()}: Initializes from HuggingFace config + +\item \code{qwen3_from_pretrained()}: Initializes and loads pretrained weights from HF Hub + +}} diff --git a/tests/testthat/test-ministral.R b/tests/testthat/test-ministral.R new file mode 100644 index 0000000..c207d16 --- /dev/null +++ b/tests/testthat/test-ministral.R @@ -0,0 +1,103 @@ +test_that("Can create a ministral model from Mistral-7B", { + + skip_on_ci() # too large for CI runners + + model <- ministral_from_pretrained("mistralai/Mistral-7B-v0.1") + + model$to(dtype = torch_float32()) + model$eval() + + # Input tokens (R uses 1-based embedding indexing, so add 1 to Python indices) + # This corresponds to Python input_ids = [2, 3, 4, 5, 6] + with_no_grad({ + pred <- model(torch_tensor(matrix(c(3L, 4L, 5L, 6L, 7L), nrow = 1))) + }) + + # Reference from Python transformers MistralForCausalLM + # output[0, 0, :5] for input_ids [[2, 3, 4, 5, 6]] + out <- pred[1, 1, 1:5] + reference <- c(1.4050, 7.4602, -0.5856, 0.5367, 1.4043) + expect_equal(as.numeric(out), reference, tolerance = 1e-4) +}) + +test_that("Can create a ministral model with custom config", { + model <- ministral( + vocab_size = 1000, + n_embd = 256, + n_inter = 512, + n_head = 8, + n_kv_head = 2, + head_dim = 32, + n_layer = 2, + max_pos = 512 + ) + + input_ids <- torch_randint(1, 1000, c(2, 10), dtype = torch_long()) + with_no_grad({ + out <- model(input_ids) + }) + + expect_equal(out$shape, c(2, 10, 1000)) +}) + +test_that("Can generate text with ministral", { + skip_on_ci() # too large for CI runners + skip_if_not_installed("tok") + + model <- ministral_from_pretrained("mistralai/Mistral-7B-v0.1") + model$eval() + + tokenizer <- tok::tokenizer$from_pretrained("mistralai/Mistral-7B-v0.1") + + # Generation parameters + prompt <- "The capital of France is" + max_new_tokens <- 20 + temperature <- 0.7 + top_k <- 50 + eos_token_id <- 2 + + # Encode prompt (tok returns 0-indexed token ids) + encoded <- tokenizer$encode(prompt) + # Add 1 for R's 1-based embedding indexing + idx <- torch_tensor(encoded$ids, dtype = torch_long())$unsqueeze(1) + 1L + + generated_ids <- c() + prev_text <- "" + + cat("\n", prompt, sep = "") + for (i in seq_len(max_new_tokens)) { + with_no_grad({ + logits <- model(idx) + }) + + # Get logits for last position + logits <- logits[1, -1, ] + + # Apply temperature and top-k + logits <- logits / temperature + c(topk_values, topk_indices) %<-% logits$topk(top_k) + logits <- torch_full_like(logits, -Inf) + logits$scatter_(1, topk_indices, topk_values) + + # Sample + probs <- nnf_softmax(logits, dim = -1) + id_next <- torch_multinomial(probs, num_samples = 1L) + id_next_val <- as.integer(id_next) - 1L + + if (id_next_val == eos_token_id) break + + generated_ids <- c(generated_ids, id_next_val) + + # Stream output: decode full sequence and print only new characters + full_text <- tokenizer$decode(generated_ids) + new_text <- substr(full_text, nchar(prev_text) + 1, nchar(full_text)) + cat(new_text) + prev_text <- full_text + + idx <- torch_cat(list(idx, id_next$unsqueeze(1)), dim = 2) + } + cat("\n") + + # Basic check - we generated something + expect_true(length(generated_ids) > 0) +}) diff --git a/tests/testthat/test-qwen3.R b/tests/testthat/test-qwen3.R new file mode 100644 index 0000000..14e59eb --- /dev/null +++ b/tests/testthat/test-qwen3.R @@ -0,0 +1,47 @@ +test_that("Can create a qwen3 model from Qwen3.5-0.8B", { + + skip_on_ci() # too large for CI runners + + model <- qwen3_from_pretrained("Qwen/Qwen3.5-0.8B") + + model$to(dtype = torch_float32()) + model$eval() + + # Input tokens: R uses 1-based indexing, so add 1 to Python indices + # This corresponds to Python input_ids = [1, 2, 3, 4, 5] + with_no_grad({ + pred <- model(torch_tensor(matrix(2:6, nrow = 1), dtype = torch_long())) + }) + + # Reference from Python transformers Qwen3_5ForCausalLM (float32) + # output[0, 0, :5] for input_ids [[1, 2, 3, 4, 5]] + out <- pred[1, 1, 1:5] + reference <- c(3.3154, 6.2008, 4.0241, 2.3807, 1.2581) + expect_equal(as.numeric(out), reference, tolerance = 1e-3) +}) + +test_that("Can create a qwen3 model with custom config", { + model <- qwen3( + vocab_size = 1000, + n_embd = 64, + n_inter = 128, + n_head = 2, + n_kv_head = 1, + head_dim = 32, + n_layer = 4, + max_pos = 1024, + layer_types = c("linear_attention", "linear_attention", + "linear_attention", "full_attention"), + n_k_heads = 4, + n_v_heads = 4, + k_head_dim = 16, + v_head_dim = 16 + ) + + input_ids <- torch_randint(1, 1000, c(2, 10), dtype = torch_long()) + with_no_grad({ + out <- model(input_ids) + }) + + expect_equal(out$shape, c(2, 10, 1000)) +})