You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

MPSGraph变量赋值操作为何无法更新变量值?

MPSGraph调用assignVariable后变量值未更新的问题分析

问题描述

我正在学习MPSGraph API,以下代码在第二次调用run方法后,var变量的值并未改变,请问这是什么原因?

代码示例

#import <Foundation/Foundation.h>
#import <MetalPerformanceShadersGraph/MetalPerformanceShadersGraph.h>

void run(void);

void run(void) {
    MPSGraph *graph = [MPSGraph new];

    float test1 = 2.0;
    float test2 = 3.0;
    MPSGraphTensor *constant = [graph constantWithScalar:test1
                                                dataType:MPSDataTypeFloat32];
    MPSGraphTensor *var = [graph variableWithData:[NSData dataWithBytes:&test2 length:sizeof(test2)]
                                            shape:@[@1]
                                         dataType:MPSDataTypeFloat32
                                             name:@"var"];

    MPSGraphTensorDataDictionary *result = [graph runWithFeeds:@@{}
                                                 targetTensors:@[constant, var]
                                              targetOperations:NULL];

    float test3;
    NSInteger temp = sizeof(test3);
    [result[var].mpsndarray readBytes:&test3
                          strideBytes:&temp];
    NSLog(@"%f", test3);

    [result[constant].mpsndarray readBytes:&test3
                               strideBytes:&temp];
    NSLog(@"%f", test3);
    
    MPSGraphOperation *op = [graph assignVariable:var
                                withValueOfTensor:constant
                                             name:NULL];
    
    result = [graph runWithFeeds:@@{}
                   targetTensors:@[var]
                targetOperations:@[op]];

    [result[var].mpsndarray readBytes:&test3
                          strideBytes:&temp];
    
    NSLog(@"%f", test3);
}

int main(int argc, const char * argv[]) {
    @autoreleasepool {
        run();
    }
    return 0;
}

输出结果

2022-09-27 13:26:35.708641-0400 MPSGraphExample[9643:2669233] Metal API Validation Enabled
2022-09-27 13:26:35.732058-0400 MPSGraphExample[9643:2669233] 3.000000
2022-09-27 13:26:35.732097-0400 MPSGraphExample[9643:2669233] 2.000000
2022-09-27 13:26:35.733821-0400 MPSGraphExample[9643:2669233] 3.000000
Program ended with exit code: 0

原因分析

第二次调用run时,你同时指定了targetTensors:@[var]和targetOperations:@[op],但MPSGraph的执行逻辑是先计算所有targetTensors的结果,再执行targetOperations中的操作。这就导致你读取到的var值是assign操作执行前的初始值,而非赋值后的结果。

此外,assignVariable:方法会返回一个代表赋值后变量状态的MPSGraphTensor,但你没有利用这个返回值获取更新后的结果,而是直接读取原变量节点,这也是问题的核心之一。

解决方法

方法一:使用assign操作返回的Tensor获取结果

修改代码,将assignVariable:返回的Tensor作为targetTensors的目标,确保获取的是赋值后的结果:

// 保存assign操作返回的Tensor,该Tensor代表赋值后的变量状态
MPSGraphTensor *updatedVar = [graph assignVariable:var
                                withValueOfTensor:constant
                                             name:NULL];

// 第二次run时,以updatedVar为目标Tensor
result = [graph runWithFeeds:@@{}
               targetTensors:@[updatedVar]
            targetOperations:@[op]];

方法二:分两次执行run操作

先单独执行assign操作更新变量状态,再读取变量的当前值:

// 先执行assign操作,更新变量状态
[graph runWithFeeds:@@{}
       targetTensors:nil
    targetOperations:@[op]];

// 再读取变量的更新后的值
result = [graph runWithFeeds:@@{}
               targetTensors:@[var]
            targetOperations:nil];

修改后,第二次输出的test3会变为2.000000,符合预期。


内容的提问来源于stack exchange,提问作者user1139069

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.18 05:50:56