Current section
Files
Jump to
Current section
Files
examples/structured_regularizers_live.exs
# Structured Regularizers - Live API Example
#
# Demonstrates custom loss computation with composable regularizers
# using the live Tinker API.
#
# Run with: TINKER_API_KEY=your-key mix run examples/structured_regularizers_live.exs
#
# Environment Variables:
# TINKER_API_KEY (required) - API authentication key
# TINKER_BASE_URL (optional) - API endpoint URL
# TINKER_BASE_MODEL (optional) - Model identifier, defaults to Llama-3.1-8B
defmodule Tinkex.Examples.StructuredRegularizersLive do
@default_base_url "https://tinker.thinkingmachines.dev/services/tinker-prod"
@default_model "meta-llama/Llama-3.1-8B"
@await_timeout :infinity
alias Tinkex.Error
alias Tinkex.Types.TensorData
def run do
{:ok, _} = Application.ensure_all_started(:tinkex)
api_key = fetch_env!("TINKER_API_KEY")
base_url = System.get_env("TINKER_BASE_URL", @default_base_url)
base_model = System.get_env("TINKER_BASE_MODEL", @default_model)
IO.puts("""
================================================================================
Structured Regularizers - Live API Example
================================================================================
Base URL: #{base_url}
Model: #{base_model}
""")
config = Tinkex.Config.new(api_key: api_key, base_url: base_url)
with {:ok, service} <- Tinkex.ServiceClient.start_link(config: config),
{:ok, training} <- create_training_client(service, base_model),
{:ok, datum} <- build_training_datum(training, base_model) do
run_custom_loss_with_regularizers(training, datum)
else
{:error, %Error{} = error} ->
halt_with_error("Initialization failed", error)
{:error, other} ->
halt("Initialization failed: #{inspect(other)}")
end
end
defp create_training_client(service, base_model) do
IO.puts("Creating training client...")
Tinkex.ServiceClient.create_lora_training_client(service, base_model,
lora_config: %Tinkex.Types.LoraConfig{rank: 16}
)
end
defp build_training_datum(training, base_model) do
prompt = "The quick brown fox jumps over the lazy dog."
IO.puts("Building training datum from prompt: #{prompt}")
case Tinkex.Types.ModelInput.from_text(prompt,
model_name: base_model,
training_client: training
) do
{:ok, model_input} ->
target_tokens = first_chunk_tokens(model_input)
IO.puts("Token count: #{length(target_tokens)}")
datum =
Tinkex.Types.Datum.new(%{
model_input: model_input,
loss_fn_inputs: %{
target_tokens: to_tensor(target_tokens, :int64),
weights: to_tensor(List.duplicate(1.0, length(target_tokens)), :float32)
}
})
{:ok, datum}
{:error, _} = error ->
error
end
end
defp run_custom_loss_with_regularizers(training, datum) do
IO.puts("\n--- Defining Custom Loss + Regularizers ---\n")
loss_fn = fn _data, [logprobs] ->
base = Nx.negate(Nx.mean(logprobs))
vocab_size = Nx.axis_size(logprobs, -1)
uniform = Nx.divide(Nx.tensor(1.0, type: Nx.type(logprobs)), vocab_size)
reference_logprobs = uniform |> Nx.broadcast(Nx.shape(logprobs)) |> Nx.log()
pair_logprobs = rotate_last_axis(logprobs)
l1 = NxPenalties.Penalties.l1(logprobs, reduction: :mean)
l2 = NxPenalties.Penalties.l2(logprobs, reduction: :mean, center: :mean)
elastic = NxPenalties.Penalties.elastic_net(logprobs, l1_ratio: 0.6, reduction: :mean)
entropy =
NxPenalties.Divergences.entropy(logprobs,
mode: :bonus,
reduction: :mean,
temperature: 0.5
)
kl_forward =
NxPenalties.Divergences.kl_divergence(logprobs, reference_logprobs,
reduction: :mean,
direction: :forward
)
kl_reverse =
NxPenalties.Divergences.kl_divergence(logprobs, reference_logprobs,
reduction: :mean,
direction: :reverse
)
kl_symmetric =
NxPenalties.Divergences.kl_divergence(logprobs, reference_logprobs,
reduction: :mean,
symmetric: true
)
consistency =
NxPenalties.Constraints.consistency(logprobs, pair_logprobs,
metric: :mse,
reduction: :mean
)
orthogonality = NxPenalties.Constraints.orthogonality(logprobs, mode: :soft)
loss_fn_for_grad = fn lp -> Nx.sum(lp) end
gradient_penalty =
NxPenalties.GradientPenalty.gradient_penalty(loss_fn_for_grad, logprobs, target_norm: 1.0)
total =
base
|> Nx.add(Nx.multiply(l1, 0.01))
|> Nx.add(Nx.multiply(l2, 0.005))
|> Nx.add(Nx.multiply(elastic, 0.002))
|> Nx.add(Nx.multiply(entropy, 0.001))
|> Nx.add(Nx.multiply(kl_forward, 0.01))
|> Nx.add(Nx.multiply(kl_reverse, 0.01))
|> Nx.add(Nx.multiply(kl_symmetric, 0.005))
|> Nx.add(Nx.multiply(consistency, 0.02))
|> Nx.add(Nx.multiply(orthogonality, 0.003))
|> Nx.add(Nx.multiply(gradient_penalty, 0.001))
metrics =
if tracing?(logprobs) do
%{}
else
%{
"base_nll" => Nx.to_number(base),
"l1" => Nx.to_number(l1),
"l2" => Nx.to_number(l2),
"elastic_net" => Nx.to_number(elastic),
"entropy" => Nx.to_number(entropy),
"kl_forward" => Nx.to_number(kl_forward),
"kl_reverse" => Nx.to_number(kl_reverse),
"kl_symmetric" => Nx.to_number(kl_symmetric),
"consistency" => Nx.to_number(consistency),
"orthogonality" => Nx.to_number(orthogonality),
"gradient_penalty" => Nx.to_number(gradient_penalty),
"custom_perplexity" => Nx.to_number(Nx.exp(base))
}
end
{total, metrics}
end
IO.puts("\n--- Running forward_backward_custom (Live API) ---\n")
start_time = System.monotonic_time(:millisecond)
{:ok, task} = Tinkex.TrainingClient.forward_backward_custom(training, [datum], loss_fn)
case Task.await(task, @await_timeout) do
{:ok, output} ->
duration_ms = System.monotonic_time(:millisecond) - start_time
display_results(output, duration_ms)
{:error, %Error{} = error} ->
halt_with_error("Custom loss computation failed", error)
{:error, other} ->
halt("Custom loss computation failed: #{inspect(other)}")
end
end
defp rotate_last_axis(tensor) do
size = Nx.axis_size(tensor, -1)
# Simple rotate-right by 1 along last axis using slice/concat (Nx 0.9 safe)
tail = Nx.slice_along_axis(tensor, size - 1, 1, axis: -1)
head = Nx.slice_along_axis(tensor, 0, size - 1, axis: -1)
Nx.concatenate([tail, head], axis: -1)
end
defp display_results(output, duration_ms) do
IO.puts("Completed in #{duration_ms}ms\n")
IO.puts("=== Metrics ===")
Enum.each(output.metrics, fn {k, v} -> IO.puts("#{k}: #{Float.round(v, 6)}") end)
IO.puts("\n================================================================================")
IO.puts("Success! Custom loss with regularizer terms computed via live Tinker API.")
IO.puts("================================================================================")
end
# Helpers
defp fetch_env!(var) do
case System.get_env(var) do
nil -> halt("Set #{var} environment variable to run this example")
value -> value
end
end
defp halt_with_error(prefix, %Error{} = error) do
IO.puts(:stderr, "#{prefix}: #{Error.format(error)}")
if error.data, do: IO.puts(:stderr, "Error data: #{inspect(error.data)}")
System.halt(1)
end
defp halt(message) do
IO.puts(:stderr, message)
System.halt(1)
end
defp first_chunk_tokens(%Tinkex.Types.ModelInput{chunks: [chunk | _]}) do
Map.get(chunk, :tokens) || Map.get(chunk, "tokens") || []
end
defp first_chunk_tokens(_), do: []
defp to_tensor(tokens, dtype) when is_list(tokens) do
%TensorData{data: tokens, dtype: dtype, shape: [length(tokens)]}
end
defp to_tensor(_, dtype), do: %TensorData{data: [], dtype: dtype, shape: [0]}
defp tracing?(%Nx.Tensor{data: %Nx.Defn.Expr{}}), do: true
defp tracing?(_), do: false
end
Tinkex.Examples.StructuredRegularizersLive.run()