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 120_000
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))
l1 = Nx.multiply(0.01, Nx.sum(Nx.abs(logprobs)))
probs = Nx.exp(logprobs)
entropy = Nx.sum(Nx.multiply(probs, logprobs))
entropy_term = Nx.multiply(0.001, entropy)
total =
base
|> Nx.add(l1)
|> Nx.add(entropy_term)
metrics = %{
"base_nll" => base,
"l1" => l1,
"entropy" => entropy_term,
"custom_perplexity" => Nx.exp(base)
}
{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 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]}
end
Tinkex.Examples.StructuredRegularizersLive.run()