Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 39 additions & 0 deletions LLama.Unittest/ModelsParamsTests.cs
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
using LLama.Common;
using System.Text.Json;
using LLama.Abstractions;
using LLama.Exceptions;
using LLama.Extensions;

namespace LLama.Unittest
{
Expand All @@ -20,6 +22,7 @@ public void SerializeRoundTripSystemTextJson()
ContextSize = 42,
GpuLayerCount = 111,
TensorSplits = { [0] = 3 },
Devices = { "Vulkan1", "CPU" },
MetadataOverrides =
{
new MetadataOverride("hello", true),
Expand All @@ -46,12 +49,48 @@ public void SerializeRoundTripSystemTextJson()
actual.TensorBufferOverrides = null!;
expected.TensorBufferOverrides = null!;

// Same deal
Assert.True(expected.Devices.SequenceEqual(actual.Devices));
actual.Devices = null!;
expected.Devices = null!;

// Check encoding is the same
var b1 = expected.Encoding.GetBytes("Hello");
var b2 = actual.Encoding.GetBytes("Hello");
Assert.True(b1.SequenceEqual(b2));

Assert.Equal(expected, actual);
}

[Fact]
public void UnknownDeviceThrows()
{
var @params = new ModelParams("abc/123")
{
Devices = { "NoSuchDevice" },
};

var ex = Assert.Throws<UnknownDeviceException>(() => @params.ToLlamaModelParams(out _));

Assert.Equal("NoSuchDevice", ex.RequestedDevice);
Assert.Contains("CPU", ex.AvailableDevices);
Assert.Contains("NoSuchDevice", ex.Message);
Assert.Contains("CPU", ex.Message);
}

[Fact]
public unsafe void DeviceNameMatchIsCaseInsensitive()
{
var @params = new ModelParams("abc/123")
{
Devices = { "cpu" },
};

using var disposer = @params.ToLlamaModelParams(out var result);

Assert.True(result.devices != null);
Assert.NotEqual(IntPtr.Zero, result.devices[0]);
Assert.Equal(IntPtr.Zero, result.devices[1]);
}
}
}
3 changes: 3 additions & 0 deletions LLama.Web/Common/ModelOptions.cs
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,9 @@ public class ModelOptions
/// <inheritdoc />
public List<TensorBufferOverride> TensorBufferOverrides { get; set; } = new();

/// <inheritdoc />
public List<string> Devices { get; set; } = new();

/// <inheritdoc />
public int GpuLayerCount { get; set; } = 20;

Expand Down
11 changes: 11 additions & 0 deletions LLama/Abstractions/IModelParams.cs
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,17 @@ public interface IModelParams
/// </summary>
List<TensorBufferOverride> TensorBufferOverrides { get; }

/// <summary>
/// Names of the backend devices the model may use, in priority order, e.g. "Vulkan1" or "CUDA0" (see <see cref="NativeApi.ggml_backend_dev_name"/>).
/// Equivalent to --device on the llama.cpp command line or <c>devices</c> in <c>llama_model_params</c>.
/// </summary>
/// <remarks>
/// When empty, llama.cpp picks the devices itself: all discrete GPUs, or the first integrated GPU when there is no discrete GPU.
/// Setting this list explicitly bypasses that selection, which is the only way to run on an integrated GPU in a machine that also has a discrete GPU.
/// <see cref="MainGpu"/> is an index into this list when it is non-empty. Names are matched case insensitively; a name that does not match any available device throws <see cref="Exceptions.UnknownDeviceException"/> when the model is loaded.
/// </remarks>
List<string> Devices { get; }

/// <summary>
/// Number of layers to run in VRAM / GPU memory (n_gpu_layers)
/// </summary>
Expand Down
3 changes: 3 additions & 0 deletions LLama/Common/ModelParams.cs
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,9 @@ public record ModelParams
/// <inheritdoc />
public List<TensorBufferOverride> TensorBufferOverrides { get; set; } = new();

/// <inheritdoc />
public List<string> Devices { get; set; } = new();

/// <inheritdoc />
public int GpuLayerCount { get; set; } = 20;

Expand Down
33 changes: 33 additions & 0 deletions LLama/Exceptions/UnknownDeviceException.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
using System;
using System.Collections.Generic;

namespace LLama.Exceptions;

/// <summary>
/// Thrown when a device name in <see cref="Abstractions.IModelParams.Devices"/> does not match any available ggml backend device
/// </summary>
public class UnknownDeviceException
: Exception
{
/// <summary>
/// The device name which was requested but could not be found
/// </summary>
public string RequestedDevice { get; }

/// <summary>
/// Names of all devices available on this machine, as returned by <see cref="Native.NativeApi.ggml_backend_dev_name"/>
/// </summary>
public IReadOnlyList<string> AvailableDevices { get; }

/// <summary>
/// Create a new UnknownDeviceException
/// </summary>
/// <param name="requestedDevice">The device name which was requested but could not be found</param>
/// <param name="availableDevices">Names of all devices available on this machine</param>
public UnknownDeviceException(string requestedDevice, IReadOnlyList<string> availableDevices)
: base($"Unknown device '{requestedDevice}'. Available devices: {string.Join(", ", availableDevices)}")
{
RequestedDevice = requestedDevice;
AvailableDevices = availableDevices;
}
}
64 changes: 64 additions & 0 deletions LLama/Extensions/IModelParamsExtensions.cs
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
using System;
using System.Text;
using LLama.Abstractions;
using LLama.Exceptions;
using LLama.Native;
using System.Collections.Generic;

Expand All @@ -20,6 +21,7 @@
/// <returns></returns>
/// <exception cref="FileNotFoundException"></exception>
/// <exception cref="ArgumentException"></exception>
/// <exception cref="UnknownDeviceException">Thrown if a name in <see cref="IModelParams.Devices"/> does not match any available device</exception>
public static IDisposable ToLlamaModelParams(this IModelParams @params, out LLamaModelParams result)
{
var supportsMmap = NativeApi.llama_supports_mmap();
Expand Down Expand Up @@ -56,6 +58,12 @@
result.tensor_buft_overrides = ConvertOverrides(@params.TensorBufferOverrides, disposer);
}

// Add device list
unsafe
{
result.devices = ConvertDevices(@params.Devices, disposer);
}

// Add metadata overrides
if (@params.MetadataOverrides.Count == 0)
{
Expand Down Expand Up @@ -123,12 +131,68 @@
if (string.IsNullOrEmpty(name))
continue;

result[name] = buft;

Check warning on line 134 in LLama/Extensions/IModelParamsExtensions.cs

View workflow job for this annotation

GitHub Actions / macOS ARM64 Metal

Possible null reference argument for parameter 'key' in 'IntPtr Dictionary<string, IntPtr>.this[string key]'.

Check warning on line 134 in LLama/Extensions/IModelParamsExtensions.cs

View workflow job for this annotation

GitHub Actions / Linux ARM64 CPU

Possible null reference argument for parameter 'key' in 'IntPtr Dictionary<string, IntPtr>.this[string key]'.

Check warning on line 134 in LLama/Extensions/IModelParamsExtensions.cs

View workflow job for this annotation

GitHub Actions / Windows x64 CPU

Possible null reference argument for parameter 'key' in 'IntPtr Dictionary<string, IntPtr>.this[string key]'.

Check warning on line 134 in LLama/Extensions/IModelParamsExtensions.cs

View workflow job for this annotation

GitHub Actions / Linux x64 CPU

Possible null reference argument for parameter 'key' in 'IntPtr Dictionary<string, IntPtr>.this[string key]'.
}

return result;
}

private static unsafe IntPtr* ConvertDevices(List<string> deviceNames, GroupDisposable disposer)
{
// Early out if no devices were requested (llama.cpp will choose its own)
if (deviceNames.Count == 0)
return null;

// Map device name -> ggml_backend_dev_t, keeping the names in native order for error messages
var devicesByName = new Dictionary<string, IntPtr>();
var availableNames = new List<string>();
var deviceCount = NativeApi.ggml_backend_dev_count();
for (nuint i = 0; i < deviceCount; i++)
{
var device = NativeApi.ggml_backend_dev_get(i);
if (device == IntPtr.Zero)
continue;

var name = NativeApi.ggml_backend_dev_name(device).PtrToString();
if (string.IsNullOrEmpty(name))
continue;

devicesByName[name!] = device;
availableNames.Add(name!);
}

// Resolve the requested names in order. One extra slot for the null terminator.
var devicesArray = new IntPtr[deviceNames.Count + 1];
for (var i = 0; i < deviceNames.Count; i++)
devicesArray[i] = ResolveDevice(deviceNames[i], devicesByName, availableNames);

// Pin the array so it can be passed to native code
var pin = devicesArray.AsMemory().Pin();
disposer.Add(pin);
return (IntPtr*)pin.Pointer;
}

/// <summary>
/// Find the device with the given name. Tries an exact match first, then a case insensitive match.
/// </summary>
/// <exception cref="UnknownDeviceException">Thrown if no device matches the name</exception>
private static IntPtr ResolveDevice(string name, Dictionary<string, IntPtr> devicesByName, List<string> availableNames)
{
if (!string.IsNullOrEmpty(name))
{
if (devicesByName.TryGetValue(name, out var device))
return device;

foreach (var pair in devicesByName)
{
if (string.Equals(pair.Key, name, StringComparison.OrdinalIgnoreCase))
return pair.Value;
}
}

throw new UnknownDeviceException(name ?? "", availableNames);
}

private static unsafe LLamaModelTensorBufferOverride* ConvertOverrides(List<TensorBufferOverride> overrides, GroupDisposable disposer)
{
// Early out if there are no overrides
Expand Down
6 changes: 3 additions & 3 deletions LLama/Native/LLamaModelParams.cs
Original file line number Diff line number Diff line change
Expand Up @@ -9,10 +9,10 @@ namespace LLama.Native
public unsafe struct LLamaModelParams
{
/// <summary>
/// NULL-terminated list of devices to use for offloading (if NULL, all available devices are used)
/// todo: add support for llama_model_params.devices
/// NULL-terminated list of devices to use for offloading (if NULL, all available devices are used).
/// Each element is a <c>ggml_backend_dev_t</c> as returned by <see cref="NativeApi.ggml_backend_dev_get"/>.
/// </summary>
private IntPtr devices;
public IntPtr* devices;

/// <summary>
/// NULL-terminated list of buffer types to use for tensors that match a pattern
Expand Down
24 changes: 22 additions & 2 deletions LLama/Native/Load/NativeLibraryUtils.cs
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,18 @@ namespace LLama.Native
{
internal static class NativeLibraryUtils
{
/// <summary>
/// Handle of the ggml library loaded as a dependency by <see cref="TryLoadLibrary"/>, or IntPtr.Zero if it has not been loaded.
/// Used by the DllImport resolver so that P/Invokes into ggml reuse this library instead of relying on the system search path.
/// </summary>
internal static IntPtr LoadedGgmlHandle;

/// <summary>
/// Handle of the ggml-base library loaded as a dependency by <see cref="TryLoadLibrary"/>, or IntPtr.Zero if it has not been loaded.
/// Used by the DllImport resolver so that P/Invokes into ggml-base reuse this library instead of relying on the system search path.
/// </summary>
internal static IntPtr LoadedGgmlBaseHandle;

/// <summary>
/// Try to load libllama/mtmd, using CPU feature detection to try and load a more specialised DLL if possible
/// </summary>
Expand Down Expand Up @@ -64,7 +76,8 @@ internal static IntPtr TryLoadLibrary(NativeLibraryConfig config, out INativeLib
var dependencyPaths = new List<string>();

// We should always load ggml-base from the current runtime directory
dependencyPaths.Add(Path.Combine(currentRuntimeDirectory, $"{libPrefix}ggml-base{ext}"));
var ggmlBasePath = Path.Combine(currentRuntimeDirectory, $"{libPrefix}ggml-base{ext}");
dependencyPaths.Add(ggmlBasePath);

// If the library has metadata, we can check if we need to load additional dependencies
if (library.Metadata != null)
Expand Down Expand Up @@ -114,13 +127,20 @@ internal static IntPtr TryLoadLibrary(NativeLibraryConfig config, out INativeLib
}

// And finally, we can add ggml
dependencyPaths.Add(Path.Combine(currentRuntimeDirectory, $"{libPrefix}ggml{ext}"));
var ggmlPath = Path.Combine(currentRuntimeDirectory, $"{libPrefix}ggml{ext}");
dependencyPaths.Add(ggmlPath);

// Now, we will loop through our dependencyPaths and try to load them one by one
foreach (var dependencyPath in dependencyPaths)
{
// Try to load the dependency
var dependencyResult = TryLoad(dependencyPath, description.SearchDirectories, config.LogCallback);

// Keep the ggml/ggml-base handles so the DllImport resolver can hand them out for P/Invokes into those libraries
if (dependencyPath == ggmlBasePath)
LoadedGgmlBaseHandle = dependencyResult;
else if (dependencyPath == ggmlPath)
LoadedGgmlHandle = dependencyResult;

// If we successfully loaded the library, log it
if (dependencyResult != IntPtr.Zero)
Expand Down
11 changes: 11 additions & 0 deletions LLama/Native/NativeApi.Load.cs
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,17 @@ private static void SetDllImportResolver()
return _loadedLlamaHandle;
}

if (name == ggmlLibraryName || name == ggmlBaseLibraryName)
{
// ggml and ggml-base are loaded by full path as dependencies of llama, so the default runtime
// resolution cannot be relied on to find them by name (it fails on macOS). Make sure llama (and
// therefore its dependencies) has been loaded, then return the dependency handle.
if (_loadedLlamaHandle == IntPtr.Zero)
_loadedLlamaHandle = NativeLibraryUtils.TryLoadLibrary(NativeLibraryConfig.LLama, out _loadedLLamaLibrary);

return name == ggmlLibraryName ? NativeLibraryUtils.LoadedGgmlHandle : NativeLibraryUtils.LoadedGgmlBaseHandle;
}

if (name == "mtmd")
{
// If we've already loaded Mtmd return the handle that was loaded last time.
Expand Down
8 changes: 8 additions & 0 deletions LLama/Native/NativeApi.cs
Original file line number Diff line number Diff line change
Expand Up @@ -448,6 +448,14 @@ public static string llama_split_prefix(string splitPath, int splitNo, int split
[DllImport(ggmlLibraryName, CallingConvention = CallingConvention.Cdecl)]
public static extern IntPtr ggml_backend_dev_get(nuint i);

/// <summary>
/// Get the name of a backend device (e.g. "CPU", "Vulkan0", "CUDA1")
/// </summary>
/// <param name="dev">Backend device pointer</param>
/// <returns>Pointer to a null terminated UTF-8 string, owned by the device</returns>
[DllImport(ggmlBaseLibraryName, CallingConvention = CallingConvention.Cdecl)]
public static extern IntPtr ggml_backend_dev_name(IntPtr dev);

/// <summary>
/// Get the buffer type for a backend device
/// </summary>
Expand Down
Loading