Skip to content

Commit 621c32d

Browse files
committed
feat: support configuring init container image
Signed-off-by: rudeigerc <rudeigerc@gmail.com>
1 parent 768bdc1 commit 621c32d

7 files changed

Lines changed: 50 additions & 39 deletions

File tree

pkg/controller/inference/service_controller.go

Lines changed: 12 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@ import (
4444

4545
coreapi "github.com/inftyai/llmaz/api/core/v1alpha1"
4646
inferenceapi "github.com/inftyai/llmaz/api/inference/v1alpha1"
47+
"github.com/inftyai/llmaz/pkg"
4748
helper "github.com/inftyai/llmaz/pkg/controller_helper"
4849
modelSource "github.com/inftyai/llmaz/pkg/controller_helper/modelsource"
4950
"github.com/inftyai/llmaz/pkg/util"
@@ -116,7 +117,12 @@ func (r *ServiceReconciler) Reconcile(ctx context.Context, req ctrl.Request) (ct
116117
return ctrl.Result{}, err
117118
}
118119

119-
workloadApplyConfiguration := buildWorkloadApplyConfiguration(service, models)
120+
initContainerImage := configs.InitContainerImage
121+
if initContainerImage == "" {
122+
initContainerImage = pkg.LOADER_IMAGE
123+
}
124+
125+
workloadApplyConfiguration := buildWorkloadApplyConfiguration(service, models, initContainerImage)
120126
if err := setControllerReferenceForWorkload(service, workloadApplyConfiguration, r.Scheme); err != nil {
121127
return ctrl.Result{}, err
122128
}
@@ -159,7 +165,7 @@ func (r *ServiceReconciler) SetupWithManager(mgr ctrl.Manager) error {
159165
Complete(r)
160166
}
161167

162-
func buildWorkloadApplyConfiguration(service *inferenceapi.Service, models []*coreapi.OpenModel) *applyconfigurationv1.LeaderWorkerSetApplyConfiguration {
168+
func buildWorkloadApplyConfiguration(service *inferenceapi.Service, models []*coreapi.OpenModel, initContainerImage string) *applyconfigurationv1.LeaderWorkerSetApplyConfiguration {
163169
workload := applyconfigurationv1.LeaderWorkerSet(service.Name, service.Namespace)
164170

165171
leaderWorkerTemplate := applyconfigurationv1.LeaderWorkerTemplate()
@@ -169,7 +175,7 @@ func buildWorkloadApplyConfiguration(service *inferenceapi.Service, models []*co
169175
leaderWorkerTemplate.WithWorkerTemplate(service.Spec.WorkloadTemplate.WorkerTemplate)
170176

171177
// The core logic to inject additional configurations.
172-
injectModelProperties(leaderWorkerTemplate, models, service)
178+
injectModelProperties(leaderWorkerTemplate, models, service, initContainerImage)
173179

174180
spec := applyconfigurationv1.LeaderWorkerSetSpec()
175181
spec.WithLeaderWorkerTemplate(leaderWorkerTemplate)
@@ -191,15 +197,15 @@ func buildWorkloadApplyConfiguration(service *inferenceapi.Service, models []*co
191197
return workload
192198
}
193199

194-
func injectModelProperties(template *applyconfigurationv1.LeaderWorkerTemplateApplyConfiguration, models []*coreapi.OpenModel, service *inferenceapi.Service) {
200+
func injectModelProperties(template *applyconfigurationv1.LeaderWorkerTemplateApplyConfiguration, models []*coreapi.OpenModel, service *inferenceapi.Service, initContainerImage string) {
195201
isMultiNodesInference := template.LeaderTemplate != nil
196202

197203
for i, model := range models {
198204
source := modelSource.NewModelSourceProvider(model)
199205
if isMultiNodesInference {
200-
source.InjectModelLoader(template.LeaderTemplate, i)
206+
source.InjectModelLoader(template.LeaderTemplate, i, initContainerImage)
201207
}
202-
source.InjectModelLoader(template.WorkerTemplate, i)
208+
source.InjectModelLoader(template.WorkerTemplate, i, initContainerImage)
203209
}
204210

205211
// We only consider the main model's requirements for now.

pkg/controller_helper/modelsource/modelhub.go

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -22,8 +22,6 @@ import (
2222

2323
corev1 "k8s.io/api/core/v1"
2424
"k8s.io/utils/ptr"
25-
26-
"github.com/inftyai/llmaz/pkg"
2725
)
2826

2927
var _ ModelSourceProvider = &ModelHubProvider{}
@@ -57,7 +55,7 @@ func (p *ModelHubProvider) ModelPath() string {
5755
return CONTAINER_MODEL_PATH + "models--" + strings.ReplaceAll(p.modelID, "/", "--")
5856
}
5957

60-
func (p *ModelHubProvider) InjectModelLoader(template *corev1.PodTemplateSpec, index int) {
58+
func (p *ModelHubProvider) InjectModelLoader(template *corev1.PodTemplateSpec, index int, initContainerImage string) {
6159
initContainerName := MODEL_LOADER_CONTAINER_NAME
6260
if index != 0 {
6361
initContainerName += "-" + strconv.Itoa(index)
@@ -66,7 +64,7 @@ func (p *ModelHubProvider) InjectModelLoader(template *corev1.PodTemplateSpec, i
6664
// Handle initContainer.
6765
initContainer := &corev1.Container{
6866
Name: initContainerName,
69-
Image: pkg.LOADER_IMAGE,
67+
Image: initContainerImage,
7068
VolumeMounts: []corev1.VolumeMount{
7169
{
7270
Name: MODEL_VOLUME_NAME,

pkg/controller_helper/modelsource/modelsource.go

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -52,8 +52,9 @@ type ModelSourceProvider interface {
5252
ModelName() string
5353
ModelPath() string
5454
// InjectModelLoader will inject the model loader to the spec,
55-
// index refers to the suffix of the initContainer name, like model-loader, model-loader-1.
56-
InjectModelLoader(spec *corev1.PodTemplateSpec, index int)
55+
// index refers to the suffix of the initContainer name, like model-loader, model-loader-1,
56+
// initContainerImage is the image used for the model loader.
57+
InjectModelLoader(spec *corev1.PodTemplateSpec, index int, initContainerImage string)
5758
}
5859

5960
func NewModelSourceProvider(model *coreapi.OpenModel) ModelSourceProvider {

pkg/controller_helper/modelsource/modelsource_test.go

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@ import (
2323
corev1 "k8s.io/api/core/v1"
2424

2525
coreapi "github.com/inftyai/llmaz/api/core/v1alpha1"
26+
"github.com/inftyai/llmaz/pkg"
2627
"github.com/inftyai/llmaz/test/util"
2728
"github.com/inftyai/llmaz/test/util/wrapper"
2829
)
@@ -123,7 +124,7 @@ func TestEnvInjectModelLoader(t *testing.T) {
123124

124125
for _, tt := range tests {
125126
t.Run(tt.name, func(t *testing.T) {
126-
tt.provider.InjectModelLoader(tt.template, 0)
127+
tt.provider.InjectModelLoader(tt.template, 0, pkg.LOADER_IMAGE)
127128
initContainer := tt.template.Spec.InitContainers[0]
128129
assert.Subset(t, initContainer.Env, tt.template.Spec.Containers[0].Env)
129130
})

pkg/controller_helper/modelsource/uri.go

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -22,8 +22,6 @@ import (
2222

2323
corev1 "k8s.io/api/core/v1"
2424
"k8s.io/utils/ptr"
25-
26-
"github.com/inftyai/llmaz/pkg"
2725
)
2826

2927
var _ ModelSourceProvider = &URIProvider{}
@@ -73,7 +71,7 @@ func (p *URIProvider) ModelPath() string {
7371
return CONTAINER_MODEL_PATH + "models--" + splits[len(splits)-1]
7472
}
7573

76-
func (p *URIProvider) InjectModelLoader(template *corev1.PodTemplateSpec, index int) {
74+
func (p *URIProvider) InjectModelLoader(template *corev1.PodTemplateSpec, index int, initContainerImage string) {
7775
// We don't have additional operations for Ollama, just load in runtime.
7876
if p.protocol == Ollama {
7977
return
@@ -111,7 +109,7 @@ func (p *URIProvider) InjectModelLoader(template *corev1.PodTemplateSpec, index
111109
// Handle initContainer.
112110
initContainer := &corev1.Container{
113111
Name: initContainerName,
114-
Image: pkg.LOADER_IMAGE,
112+
Image: initContainerImage,
115113
VolumeMounts: []corev1.VolumeMount{
116114
{
117115
Name: MODEL_VOLUME_NAME,

test/config/others/global-configmap.yaml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6,4 +6,4 @@ metadata:
66
data:
77
config.data: |
88
scheduler-name: inftyai-scheduler
9-
init-container-image: inftyai/model-loader:v0.0.10
9+
init-container-image: docker.m.daocloud.io/inftyai/model-loader:v0.0.10

test/util/validation/validate_service.go

Lines changed: 28 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -72,14 +72,30 @@ func ValidateService(ctx context.Context, k8sClient client.Client, service *infe
7272
models = append(models, model)
7373
}
7474

75+
// Fetch global configuration.
76+
cm := corev1.ConfigMap{}
77+
if err := k8sClient.Get(ctx, types.NamespacedName{Name: "llmaz-global-config", Namespace: "llmaz-system"}, &cm); err != nil {
78+
return err
79+
}
80+
81+
data, err := helper.ParseGlobalConfigmap(&cm)
82+
if err != nil {
83+
return fmt.Errorf("failed to parse global configmap: %v", err)
84+
}
85+
86+
initContainerImage := data.InitContainerImage
87+
if initContainerImage == "" {
88+
initContainerImage = pkg.LOADER_IMAGE
89+
}
90+
7591
for index, model := range models {
7692
// Validate injecting modelLoaders
7793
if service.Spec.WorkloadTemplate.LeaderTemplate != nil {
78-
if err := ValidateModelLoader(model, index, *workload.Spec.LeaderWorkerTemplate.LeaderTemplate, service); err != nil {
94+
if err := ValidateModelLoader(model, index, *workload.Spec.LeaderWorkerTemplate.LeaderTemplate, service, initContainerImage); err != nil {
7995
return err
8096
}
8197
}
82-
if err := ValidateModelLoader(model, index, workload.Spec.LeaderWorkerTemplate.WorkerTemplate, service); err != nil {
98+
if err := ValidateModelLoader(model, index, workload.Spec.LeaderWorkerTemplate.WorkerTemplate, service, initContainerImage); err != nil {
8399
return err
84100
}
85101
}
@@ -103,15 +119,15 @@ func ValidateService(ctx context.Context, k8sClient client.Client, service *infe
103119
return err
104120
}
105121

106-
if err := ValidateConfigmap(ctx, k8sClient, service); err != nil {
122+
if err := ValidateSchedulerName(data.SchedulerName, service); err != nil {
107123
return err
108124
}
109125

110126
return nil
111127
}, util.IntegrationTimeout, util.Interval).Should(gomega.Succeed())
112128
}
113129

114-
func ValidateModelLoader(model *coreapi.OpenModel, index int, template corev1.PodTemplateSpec, service *inferenceapi.Service) error {
130+
func ValidateModelLoader(model *coreapi.OpenModel, index int, template corev1.PodTemplateSpec, service *inferenceapi.Service, initContainerImage string) error {
115131
if model.Spec.Source.URI != nil {
116132
protocol, _, _ := pkgUtil.ParseURI(string(*model.Spec.Source.URI))
117133
if protocol == modelSource.Ollama {
@@ -132,8 +148,9 @@ func ValidateModelLoader(model *coreapi.OpenModel, index int, template corev1.Po
132148
if initContainer.Name != containerName {
133149
return fmt.Errorf("unexpected initContainer name, want %s, got %s", modelSource.MODEL_LOADER_CONTAINER_NAME, initContainer.Name)
134150
}
135-
if initContainer.Image != pkg.LOADER_IMAGE {
136-
return fmt.Errorf("unexpected initContainer image, want %s, got %s", pkg.LOADER_IMAGE, initContainer.Image)
151+
152+
if initContainer.Image != initContainerImage {
153+
return fmt.Errorf("unexpected initContainer image, want %s, got %s", initContainerImage, initContainer.Image)
137154
}
138155

139156
var envStrings []string
@@ -356,25 +373,15 @@ func CheckServiceAvaliable() error {
356373
return nil
357374
}
358375

359-
func ValidateConfigmap(ctx context.Context, k8sClient client.Client, service *inferenceapi.Service) error {
360-
cm := corev1.ConfigMap{}
361-
if err := k8sClient.Get(ctx, types.NamespacedName{Name: "llmaz-global-config", Namespace: "llmaz-system"}, &cm); err != nil {
362-
return err
363-
}
364-
365-
data, err := helper.ParseGlobalConfigmap(&cm)
366-
if err != nil {
367-
return fmt.Errorf("failed to parse global configmap: %v", err)
368-
}
369-
376+
func ValidateSchedulerName(schedulerName string, service *inferenceapi.Service) error {
370377
if service.Spec.WorkloadTemplate.LeaderTemplate != nil {
371-
if service.Spec.WorkloadTemplate.LeaderTemplate.Spec.SchedulerName != data.SchedulerName {
372-
return fmt.Errorf("unexpected scheduler name %s, want %s", service.Spec.WorkloadTemplate.LeaderTemplate.Spec.SchedulerName, data.SchedulerName)
378+
if service.Spec.WorkloadTemplate.LeaderTemplate.Spec.SchedulerName != schedulerName {
379+
return fmt.Errorf("unexpected scheduler name %s, want %s", service.Spec.WorkloadTemplate.LeaderTemplate.Spec.SchedulerName, schedulerName)
373380
}
374381
}
375382

376-
if service.Spec.WorkloadTemplate.WorkerTemplate.Spec.SchedulerName != data.SchedulerName {
377-
return fmt.Errorf("unexpected scheduler name %s, want %s", service.Spec.WorkloadTemplate.WorkerTemplate.Spec.SchedulerName, data.SchedulerName)
383+
if service.Spec.WorkloadTemplate.WorkerTemplate.Spec.SchedulerName != schedulerName {
384+
return fmt.Errorf("unexpected scheduler name %s, want %s", service.Spec.WorkloadTemplate.WorkerTemplate.Spec.SchedulerName, schedulerName)
378385
}
379386

380387
return nil

0 commit comments

Comments
 (0)