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
5 changes: 4 additions & 1 deletion pkg/device-plugin/nvidiadevice/nvinternal/plugin/register.go
Original file line number Diff line number Diff line change
Expand Up @@ -110,7 +110,10 @@ func parseNvidiaNumaInfo(idx int, nvidiaTopoStr string) (int, error) {
func (plugin *NvidiaDevicePlugin) getAPIDevices() *[]*util.DeviceInfo {
devs := plugin.Devices()
klog.V(5).InfoS("getAPIDevices", "devices", devs)
nvml.Init()
if nvret := nvml.Init(); nvret != nvml.SUCCESS {
klog.Errorln("nvml Init err: ", nvret)
panic(0)
}
res := make([]*util.DeviceInfo, 0, len(devs))
for UUID := range devs {
ndev, ret := nvml.DeviceGetHandleByUUID(UUID)
Expand Down
9 changes: 7 additions & 2 deletions pkg/device-plugin/nvidiadevice/nvinternal/plugin/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -203,7 +203,12 @@ func (plugin *NvidiaDevicePlugin) Devices() rm.Devices {
func (plugin *NvidiaDevicePlugin) Start() error {
plugin.initialize()

err := plugin.Serve()
deviceNumbers, err := GetDeviceNums()
if err != nil {
return err
}

err = plugin.Serve()
if err != nil {
klog.Infof("Could not start device plugin for '%s': %s", plugin.rm.Resource(), err)
plugin.cleanup()
Expand Down Expand Up @@ -234,7 +239,7 @@ func (plugin *NvidiaDevicePlugin) Start() error {
if len(plugin.migCurrent.MigConfigs["current"]) == 1 && len(plugin.migCurrent.MigConfigs["current"][0].Devices) == 0 {
idx := 0
plugin.migCurrent.MigConfigs["current"][0].Devices = make([]int32, 0)
for idx < GetDeviceNums() {
for idx < deviceNumbers {
plugin.migCurrent.MigConfigs["current"][0].Devices = append(plugin.migCurrent.MigConfigs["current"][0].Devices, int32(idx))
idx++
}
Expand Down
21 changes: 16 additions & 5 deletions pkg/device-plugin/nvidiadevice/nvinternal/plugin/util.go
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@ package plugin
import (
"bytes"
"errors"
"fmt"
"os"
"os/exec"
"strconv"
Expand Down Expand Up @@ -92,7 +93,10 @@ func EraseNextDeviceTypeFromAnnotation(dtype string, p corev1.Pod) error {
}

func GetIndexAndTypeFromUUID(uuid string) (string, int) {
nvml.Init()
if nvret := nvml.Init(); nvret != nvml.SUCCESS {
klog.Errorln("nvml Init err: ", nvret)
panic(0)
}
originuuid := strings.Split(uuid, "[")[0]
ndev, ret := nvml.DeviceGetHandleByUUID(originuuid)
if ret != nvml.SUCCESS {
Expand Down Expand Up @@ -145,7 +149,10 @@ func GetMigUUIDFromSmiOutput(output string, uuid string, idx int) string {
}

func GetMigUUIDFromIndex(uuid string, idx int) string {
nvml.Init()
if nvret := nvml.Init(); nvret != nvml.SUCCESS {
klog.Errorln("nvml Init err: ", nvret)
panic(0)
}
originuuid := strings.Split(uuid, "[")[0]
ndev, ret := nvml.DeviceGetHandleByUUID(originuuid)
if ret != nvml.SUCCESS {
Expand Down Expand Up @@ -175,13 +182,17 @@ func GetMigUUIDFromIndex(uuid string, idx int) string {
return res
}

func GetDeviceNums() int {
nvml.Init()
func GetDeviceNums() (int, error) {
if nvret := nvml.Init(); nvret != nvml.SUCCESS {
klog.Errorln("nvml Init err: ", nvret)
return 0, fmt.Errorf("nvml Init err: %s", nvml.ErrorString(nvret))
}
count, ret := nvml.DeviceGetCount()
if ret != nvml.SUCCESS {
klog.Error(`nvml get count error ret=`, ret)
return 0, fmt.Errorf("nvml get count error ret: %s", nvml.ErrorString(ret))
}
return count
return count, nil
}

func (nv *NvidiaDevicePlugin) ApplyMigTemplate() {
Expand Down