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
相关产品推荐
相关产品推荐

